GitLab Repo

amachine.am_transformers.am_logger

  1from pathlib import Path
  2import json
  3import math
  4import torch
  5import dataclasses
  6from typing import Any
  7from .am_metrics import get_weight_norms
  8
  9_NAT_TO_BIT = 1.44269504089
 10
 11class RobustJSONEncoder(json.JSONEncoder):
 12    """Handles numpy arrays, torch tensors, and other un-serializable objects cleanly."""
 13    def default(self, o):
 14        # Catch tensors and numpy arrays
 15        if hasattr(o, 'tolist'):
 16            return o.tolist()
 17        # Catch 0-d tensors or numpy scalars
 18        if hasattr(o, 'item'):
 19            return o.item()
 20        # Fallback to string representation if all else fails
 21        try:
 22            return super().default(o)
 23        except TypeError:
 24            return str(o)
 25
 26def get_optimizer_state_metrics(optimizer: torch.optim.Optimizer) -> dict[str, float]:
 27    """Extracts aggregated metrics from the optimizer's internal state."""
 28    metrics = {}
 29    
 30    total_exp_avg_sq_norm = 0.0
 31    total_exp_avg_norm = 0.0
 32    
 33    # Track how many parameters actually have states initialized
 34    active_params = 0 
 35    
 36    for group in optimizer.param_groups:
 37        for p in group['params']:
 38            state = optimizer.state.get(p, {})
 39            
 40            # AdamW and similar optimizers use 'exp_avg' and 'exp_avg_sq'
 41            if 'exp_avg' in state and 'exp_avg_sq' in state:
 42                # Square the L2 norms to sum them correctly across the whole model
 43                total_exp_avg_norm += torch.norm(state['exp_avg'].float()).item() ** 2
 44                total_exp_avg_sq_norm += torch.norm(state['exp_avg_sq'].float()).item() ** 2
 45                active_params += 1
 46
 47    if active_params > 0:
 48        # Take the square root of the sum of squares for the global L2 norm
 49        metrics["optim/exp_avg_norm"] = math.sqrt(total_exp_avg_norm)
 50        metrics["optim/exp_avg_sq_norm"] = math.sqrt(total_exp_avg_sq_norm)
 51        
 52    return metrics
 53
 54def log_metrics(
 55    ctx,          # TrainerContext
 56    state,        # TrainerState
 57    metrics,      # WindowMetrics
 58    eval_metrics: dict[str, Any] | None = None
 59):  
 60    # Derived values
 61    perplexity = math.exp(min(metrics.mean_loss, 20))
 62    norm_str = f"{metrics.grad_norm:.3f}" if metrics.grad_norm >= 0 else "skipped"
 63    
 64    # Console output
 65    msg = (
 66        f"step {state.global_step:>7d} | loss={metrics.mean_loss:.4f} | loss_bits={metrics.mean_loss*_NAT_TO_BIT:.4f} | perplexity={perplexity:.2f} loss_ema={state.spike_loss_ema} | loss_ema (bits)={state.spike_loss_ema*_NAT_TO_BIT:.4f} |"
 67        f"grad={norm_str} | lr={metrics.lr:.2e} | tok/s={metrics.tok_per_s:.0f} | "
 68        f"ms/step={metrics.ms_per_step:.1f} | epoch={state.epoch}"
 69    )
 70    
 71    if metrics.skipped_tokens:
 72        msg += f" | skipped_tok={metrics.skipped_tokens:,}"
 73    
 74    if eval_metrics:
 75        eval_loss       = eval_metrics["eval_loss"]
 76        eval_perplexity = eval_metrics["eval_perplexity"]
 77        msg += f" | eval_loss={eval_loss:.4f} | eval_perplexity={eval_perplexity:.2f}"
 78    
 79        if "pb_max_kurtosis" in eval_metrics:
 80            msg += f" | max_kurtosis={eval_metrics['pb_max_kurtosis']:.4f}"
 81            
 82    if metrics.extra_loss_terms:
 83        extra_str = " | ".join( f"{k}={v:.4f}" for k, v in metrics.extra_loss_terms.items() if k in ( "main_loss", "aux_loss" ) )
 84        msg += f" | {extra_str}"
 85
 86
 87    print(msg)
 88    
 89    # --- Disk Logging Section ---
 90    metrics_dir = getattr(ctx.training_config, "metrics_dir", None)
 91    if metrics_dir:
 92        # Convert to Path object and create directories
 93        metrics_path = Path(metrics_dir)
 94        metrics_path.mkdir(parents=True, exist_ok=True)
 95        
 96        step_str = f"{state.global_step:07d}"
 97        
 98        # Package comprehensive train metrics
 99        train_dump = {
100            "step": state.global_step,
101            "epoch": state.epoch,
102            "derived_metrics": {
103                "perplexity": perplexity,
104                "loss_bits": metrics.mean_loss * _NAT_TO_BIT,
105                "loss_ema": state.spike_loss_ema,
106                "loss_ema_bits": (state.spike_loss_ema * _NAT_TO_BIT) if state.spike_loss_ema is not None else None,
107                "analysis_loss_ema": state.analysis_loss_ema,
108            },
109            "state": {
110                "global_tokens": state.global_tokens,
111                "spikes_skipped": state.spikes_skipped,
112                "consecutive_spikes": state.consecutive_spikes,
113            },
114            "metrics": dataclasses.asdict(metrics) 
115        }
116        
117        # Use the / operator for clean path joining
118        train_file = metrics_path / f"train_metrics_{step_str}.json"
119        with open(train_file, 'w') as f:
120            json.dump(train_dump, f, cls=RobustJSONEncoder, indent=2)
121            
122        if eval_metrics:
123            eval_dump = {
124                "step": state.global_step,
125                "epoch": state.epoch,
126                "eval_metrics": eval_metrics
127            }
128            eval_file = metrics_path / f"eval_metrics_{step_str}.json"
129            with open(eval_file, 'w') as f:
130                json.dump(eval_dump, f, cls=RobustJSONEncoder, indent=2)
131                
132    # --- Aim Tracking Section --- 
133    if ctx.aim_run is None:
134        return
135    
136    # Train payload
137    payload = {
138        "loss":               metrics.mean_loss,
139        "perplexity":         perplexity,
140        "tokens_per_s":       metrics.tok_per_s,
141        "ms_per_step":        metrics.ms_per_step,
142        "total_tokens":       state.global_tokens,
143        "epoch":              state.epoch,
144        "spikes/total":       state.spikes_skipped,
145        "spikes/rate":        state.spikes_skipped / max(state.global_step, 1),
146        "spikes/consecutive": state.consecutive_spikes,
147        "skipped_tokens":     metrics.skipped_tokens,
148    }
149    
150    if ctx.device.type == "cuda":
151        payload["vram/peak_gb"]    = torch.cuda.max_memory_allocated() / 1e9
152        payload["vram/current_gb"] = torch.cuda.memory_allocated() / 1e9
153    
154    if metrics.grad_norm >= 0:
155        avg_payload = {k.replace("grad_norm/", "grad_norm/avg/"): v for k, v in metrics.mean_block_norms.items()}
156        payload.update(avg_payload)
157    
158    if state.spike_loss_ema is not None:
159        payload["spike_loss_ema"] = state.spike_loss_ema
160    
161    if state.analysis_loss_ema is not None:
162        payload["analysis_loss_ema"] = state.analysis_loss_ema
163    
164    if metrics.extra_loss_terms:
165        payload.update({f"loss_terms/{k}": v for k, v in metrics.extra_loss_terms.items()})
166
167    # Aggregated per-block norms 
168    if metrics.per_block_norms:
169        snapshot_payload = {k.replace("grad_norm/", "grad_norm/snapshot/"): v for k, v in metrics.per_block_norms.items()}
170        payload.update(snapshot_payload)
171    
172    for i, group in enumerate(ctx.optimizer.param_groups):
173        payload[f"lr/group_{i}"] = group["lr"]
174    
175    if getattr(ctx.training_config, "track_weight_norms", False):
176        payload.update( get_weight_norms( ctx.model.named_parameters() ) )
177    
178    if getattr(ctx.training_config, "baseline_entropy_rate", None) is not None and state.analysis_loss_ema is not None:
179        loss_ema_bits = state.analysis_loss_ema * _NAT_TO_BIT
180        payload["entropy/loss_ema_bits"] = loss_ema_bits
181        payload["entropy/gap"]           = loss_ema_bits - ctx.training_config.baseline_entropy_rate
182
183    ctx.aim_run.track(payload, step=state.global_step, context={"subset": "train"})
184    
185    # Per-layer norms 
186    if metrics.per_layer_norms:
187        ctx.aim_run.track(metrics.per_layer_norms, step=state.global_step, context={"subset": "grad_detail"})
188
189    # Eval payload
190    if eval_metrics:
191
192        eval_loss       = eval_metrics["eval_loss"]
193        eval_perplexity = eval_metrics["eval_perplexity"]
194
195        eval_payload = {
196            "loss":         eval_loss,
197            "perplexity":   eval_perplexity,
198            "loss_gap":     eval_loss - metrics.mean_loss,
199            "eval_batches": eval_metrics.get("eval_batches"),
200        }
201
202        for i, group in enumerate(ctx.optimizer.param_groups):
203            eval_payload[f"lr/group_{i}"] = group["lr"]
204
205        # Dynamically capture pb_ metrics: scalars go straight into eval_payload;
206        # per-layer arrays are expanded into individual tracked points (one per
207        # layer) so Aim can plot them against layer index; strings (e.g. pb_verdict)
208        # aren't numeric and Aim's track() can't log them -- skip those.
209        for key, value in eval_metrics.items():
210            if not key.startswith("pb_"):
211                continue
212
213            if isinstance(value, bool):
214                continue  # bool is a subclass of int -- skip unless you want 0/1 logged
215
216            if isinstance(value, (int, float)):
217                eval_payload[key] = value
218
219            elif isinstance(value, (list, tuple)):
220                for layer_idx, layer_value in enumerate(value):
221                    if isinstance(layer_value, (int, float)) and not isinstance(layer_value, bool):
222                        ctx.aim_run.track(
223                            float(layer_value),
224                            name=key,
225                            step=state.global_step,
226                            context={"subset": "eval", "layer": layer_idx},
227                        )
228            # str (pb_verdict) and anything else: not trackable as a metric, skip
229
230        if getattr(ctx.training_config, "baseline_entropy_rate", None) is not None:
231            eval_loss_bits = eval_loss * _NAT_TO_BIT
232            eval_payload["entropy/loss_bits"] = eval_loss_bits
233            eval_payload["entropy/gap"] = eval_loss_bits - ctx.training_config.baseline_entropy_rate
234
235        ctx.aim_run.track(eval_payload, step=state.global_step, context={"subset": "eval"})
236
237    if ctx.device.type == "cuda":
238        torch.cuda.reset_peak_memory_stats()
class RobustJSONEncoder(json.encoder.JSONEncoder):
12class RobustJSONEncoder(json.JSONEncoder):
13    """Handles numpy arrays, torch tensors, and other un-serializable objects cleanly."""
14    def default(self, o):
15        # Catch tensors and numpy arrays
16        if hasattr(o, 'tolist'):
17            return o.tolist()
18        # Catch 0-d tensors or numpy scalars
19        if hasattr(o, 'item'):
20            return o.item()
21        # Fallback to string representation if all else fails
22        try:
23            return super().default(o)
24        except TypeError:
25            return str(o)

Handles numpy arrays, torch tensors, and other un-serializable objects cleanly.

def default(self, o):
14    def default(self, o):
15        # Catch tensors and numpy arrays
16        if hasattr(o, 'tolist'):
17            return o.tolist()
18        # Catch 0-d tensors or numpy scalars
19        if hasattr(o, 'item'):
20            return o.item()
21        # Fallback to string representation if all else fails
22        try:
23            return super().default(o)
24        except TypeError:
25            return str(o)

Implement this method in a subclass such that it returns a serializable object for o, or calls the base implementation (to raise a TypeError).

For example, to support arbitrary iterators, you could implement default like this::

def default(self, o):
    try:
        iterable = iter(o)
    except TypeError:
        pass
    else:
        return list(iterable)
    # Let the base class default method raise the TypeError
    return super().default(o)
def get_optimizer_state_metrics(optimizer: torch.optim.optimizer.Optimizer) -> dict[str, float]:
27def get_optimizer_state_metrics(optimizer: torch.optim.Optimizer) -> dict[str, float]:
28    """Extracts aggregated metrics from the optimizer's internal state."""
29    metrics = {}
30    
31    total_exp_avg_sq_norm = 0.0
32    total_exp_avg_norm = 0.0
33    
34    # Track how many parameters actually have states initialized
35    active_params = 0 
36    
37    for group in optimizer.param_groups:
38        for p in group['params']:
39            state = optimizer.state.get(p, {})
40            
41            # AdamW and similar optimizers use 'exp_avg' and 'exp_avg_sq'
42            if 'exp_avg' in state and 'exp_avg_sq' in state:
43                # Square the L2 norms to sum them correctly across the whole model
44                total_exp_avg_norm += torch.norm(state['exp_avg'].float()).item() ** 2
45                total_exp_avg_sq_norm += torch.norm(state['exp_avg_sq'].float()).item() ** 2
46                active_params += 1
47
48    if active_params > 0:
49        # Take the square root of the sum of squares for the global L2 norm
50        metrics["optim/exp_avg_norm"] = math.sqrt(total_exp_avg_norm)
51        metrics["optim/exp_avg_sq_norm"] = math.sqrt(total_exp_avg_sq_norm)
52        
53    return metrics

Extracts aggregated metrics from the optimizer's internal state.

def log_metrics( ctx, state, metrics, eval_metrics: dict[str, typing.Any] | None = None):
 55def log_metrics(
 56    ctx,          # TrainerContext
 57    state,        # TrainerState
 58    metrics,      # WindowMetrics
 59    eval_metrics: dict[str, Any] | None = None
 60):  
 61    # Derived values
 62    perplexity = math.exp(min(metrics.mean_loss, 20))
 63    norm_str = f"{metrics.grad_norm:.3f}" if metrics.grad_norm >= 0 else "skipped"
 64    
 65    # Console output
 66    msg = (
 67        f"step {state.global_step:>7d} | loss={metrics.mean_loss:.4f} | loss_bits={metrics.mean_loss*_NAT_TO_BIT:.4f} | perplexity={perplexity:.2f} loss_ema={state.spike_loss_ema} | loss_ema (bits)={state.spike_loss_ema*_NAT_TO_BIT:.4f} |"
 68        f"grad={norm_str} | lr={metrics.lr:.2e} | tok/s={metrics.tok_per_s:.0f} | "
 69        f"ms/step={metrics.ms_per_step:.1f} | epoch={state.epoch}"
 70    )
 71    
 72    if metrics.skipped_tokens:
 73        msg += f" | skipped_tok={metrics.skipped_tokens:,}"
 74    
 75    if eval_metrics:
 76        eval_loss       = eval_metrics["eval_loss"]
 77        eval_perplexity = eval_metrics["eval_perplexity"]
 78        msg += f" | eval_loss={eval_loss:.4f} | eval_perplexity={eval_perplexity:.2f}"
 79    
 80        if "pb_max_kurtosis" in eval_metrics:
 81            msg += f" | max_kurtosis={eval_metrics['pb_max_kurtosis']:.4f}"
 82            
 83    if metrics.extra_loss_terms:
 84        extra_str = " | ".join( f"{k}={v:.4f}" for k, v in metrics.extra_loss_terms.items() if k in ( "main_loss", "aux_loss" ) )
 85        msg += f" | {extra_str}"
 86
 87
 88    print(msg)
 89    
 90    # --- Disk Logging Section ---
 91    metrics_dir = getattr(ctx.training_config, "metrics_dir", None)
 92    if metrics_dir:
 93        # Convert to Path object and create directories
 94        metrics_path = Path(metrics_dir)
 95        metrics_path.mkdir(parents=True, exist_ok=True)
 96        
 97        step_str = f"{state.global_step:07d}"
 98        
 99        # Package comprehensive train metrics
100        train_dump = {
101            "step": state.global_step,
102            "epoch": state.epoch,
103            "derived_metrics": {
104                "perplexity": perplexity,
105                "loss_bits": metrics.mean_loss * _NAT_TO_BIT,
106                "loss_ema": state.spike_loss_ema,
107                "loss_ema_bits": (state.spike_loss_ema * _NAT_TO_BIT) if state.spike_loss_ema is not None else None,
108                "analysis_loss_ema": state.analysis_loss_ema,
109            },
110            "state": {
111                "global_tokens": state.global_tokens,
112                "spikes_skipped": state.spikes_skipped,
113                "consecutive_spikes": state.consecutive_spikes,
114            },
115            "metrics": dataclasses.asdict(metrics) 
116        }
117        
118        # Use the / operator for clean path joining
119        train_file = metrics_path / f"train_metrics_{step_str}.json"
120        with open(train_file, 'w') as f:
121            json.dump(train_dump, f, cls=RobustJSONEncoder, indent=2)
122            
123        if eval_metrics:
124            eval_dump = {
125                "step": state.global_step,
126                "epoch": state.epoch,
127                "eval_metrics": eval_metrics
128            }
129            eval_file = metrics_path / f"eval_metrics_{step_str}.json"
130            with open(eval_file, 'w') as f:
131                json.dump(eval_dump, f, cls=RobustJSONEncoder, indent=2)
132                
133    # --- Aim Tracking Section --- 
134    if ctx.aim_run is None:
135        return
136    
137    # Train payload
138    payload = {
139        "loss":               metrics.mean_loss,
140        "perplexity":         perplexity,
141        "tokens_per_s":       metrics.tok_per_s,
142        "ms_per_step":        metrics.ms_per_step,
143        "total_tokens":       state.global_tokens,
144        "epoch":              state.epoch,
145        "spikes/total":       state.spikes_skipped,
146        "spikes/rate":        state.spikes_skipped / max(state.global_step, 1),
147        "spikes/consecutive": state.consecutive_spikes,
148        "skipped_tokens":     metrics.skipped_tokens,
149    }
150    
151    if ctx.device.type == "cuda":
152        payload["vram/peak_gb"]    = torch.cuda.max_memory_allocated() / 1e9
153        payload["vram/current_gb"] = torch.cuda.memory_allocated() / 1e9
154    
155    if metrics.grad_norm >= 0:
156        avg_payload = {k.replace("grad_norm/", "grad_norm/avg/"): v for k, v in metrics.mean_block_norms.items()}
157        payload.update(avg_payload)
158    
159    if state.spike_loss_ema is not None:
160        payload["spike_loss_ema"] = state.spike_loss_ema
161    
162    if state.analysis_loss_ema is not None:
163        payload["analysis_loss_ema"] = state.analysis_loss_ema
164    
165    if metrics.extra_loss_terms:
166        payload.update({f"loss_terms/{k}": v for k, v in metrics.extra_loss_terms.items()})
167
168    # Aggregated per-block norms 
169    if metrics.per_block_norms:
170        snapshot_payload = {k.replace("grad_norm/", "grad_norm/snapshot/"): v for k, v in metrics.per_block_norms.items()}
171        payload.update(snapshot_payload)
172    
173    for i, group in enumerate(ctx.optimizer.param_groups):
174        payload[f"lr/group_{i}"] = group["lr"]
175    
176    if getattr(ctx.training_config, "track_weight_norms", False):
177        payload.update( get_weight_norms( ctx.model.named_parameters() ) )
178    
179    if getattr(ctx.training_config, "baseline_entropy_rate", None) is not None and state.analysis_loss_ema is not None:
180        loss_ema_bits = state.analysis_loss_ema * _NAT_TO_BIT
181        payload["entropy/loss_ema_bits"] = loss_ema_bits
182        payload["entropy/gap"]           = loss_ema_bits - ctx.training_config.baseline_entropy_rate
183
184    ctx.aim_run.track(payload, step=state.global_step, context={"subset": "train"})
185    
186    # Per-layer norms 
187    if metrics.per_layer_norms:
188        ctx.aim_run.track(metrics.per_layer_norms, step=state.global_step, context={"subset": "grad_detail"})
189
190    # Eval payload
191    if eval_metrics:
192
193        eval_loss       = eval_metrics["eval_loss"]
194        eval_perplexity = eval_metrics["eval_perplexity"]
195
196        eval_payload = {
197            "loss":         eval_loss,
198            "perplexity":   eval_perplexity,
199            "loss_gap":     eval_loss - metrics.mean_loss,
200            "eval_batches": eval_metrics.get("eval_batches"),
201        }
202
203        for i, group in enumerate(ctx.optimizer.param_groups):
204            eval_payload[f"lr/group_{i}"] = group["lr"]
205
206        # Dynamically capture pb_ metrics: scalars go straight into eval_payload;
207        # per-layer arrays are expanded into individual tracked points (one per
208        # layer) so Aim can plot them against layer index; strings (e.g. pb_verdict)
209        # aren't numeric and Aim's track() can't log them -- skip those.
210        for key, value in eval_metrics.items():
211            if not key.startswith("pb_"):
212                continue
213
214            if isinstance(value, bool):
215                continue  # bool is a subclass of int -- skip unless you want 0/1 logged
216
217            if isinstance(value, (int, float)):
218                eval_payload[key] = value
219
220            elif isinstance(value, (list, tuple)):
221                for layer_idx, layer_value in enumerate(value):
222                    if isinstance(layer_value, (int, float)) and not isinstance(layer_value, bool):
223                        ctx.aim_run.track(
224                            float(layer_value),
225                            name=key,
226                            step=state.global_step,
227                            context={"subset": "eval", "layer": layer_idx},
228                        )
229            # str (pb_verdict) and anything else: not trackable as a metric, skip
230
231        if getattr(ctx.training_config, "baseline_entropy_rate", None) is not None:
232            eval_loss_bits = eval_loss * _NAT_TO_BIT
233            eval_payload["entropy/loss_bits"] = eval_loss_bits
234            eval_payload["entropy/gap"] = eval_loss_bits - ctx.training_config.baseline_entropy_rate
235
236        ctx.aim_run.track(eval_payload, step=state.global_step, context={"subset": "eval"})
237
238    if ctx.device.type == "cuda":
239        torch.cuda.reset_peak_memory_stats()