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()