amachine.am_transformers.am_metrics
1from dataclasses import dataclass, field 2from collections import defaultdict 3 4import torch 5 6_LAYER_CONTAINERS = frozenset({"layers", "blocks", "h", "layer"}) 7 8def block_key(name: str) -> tuple[ int | None, str ]: 9 parts = name.split(".") 10 for i, part in enumerate(parts): 11 if part in _LAYER_CONTAINERS and i + 2 < len(parts) and parts[i + 1].isdigit(): 12 block = parts[i + 2] 13 component = parts[0] 14 return int(parts[i + 1]), f"{component}/{block}" 15 # No layer index — use the direct parent module as the key 16 return None, parts[-2] if len(parts) > 1 else parts[0] 17 18class GradNormTracker: 19 """Caches the parameter→block-key mapping once at init, then just reads grads each call.""" 20 21 def __init__(self, model, device): 22 self.device = device 23 # Pre-compute (param_ref, agg_key, layer_key) for every parameter once. 24 self._entries: list[tuple[torch.nn.Parameter, str, str]] = [] 25 for name, p in model.named_parameters(): 26 layer_idx, block = block_key(name) 27 agg_key = f"grad_norm/{block}" 28 layer_key = (f"grad_norm/layer_{layer_idx:03d}/{block}" 29 if layer_idx is not None 30 else f"grad_norm/model/{block}") 31 self._entries.append((p, agg_key, layer_key)) 32 33 # Pre-sort the unique key lists so the batched sync is always in a stable order. 34 seen_agg, seen_layer = {}, {} 35 for _, agg_key, layer_key in self._entries: 36 seen_agg[agg_key] = True 37 seen_layer[layer_key] = True 38 self._agg_keys = sorted(seen_agg) 39 self._layer_keys = sorted(seen_layer) 40 self._agg_idx = {k: i for i, k in enumerate(self._agg_keys)} 41 self._layer_idx = {k: i for i, k in enumerate(self._layer_keys)} 42 43 def compute(self) -> tuple[dict, dict]: 44 agg = [torch.zeros(1, device=self.device) for _ in self._agg_keys] 45 layer = [torch.zeros(1, device=self.device) for _ in self._layer_keys] 46 47 any_grad = False 48 for p, agg_key, layer_key in self._entries: 49 if p.grad is None: 50 continue 51 sq = p.grad.detach().float().pow(2).sum() 52 agg [self._agg_idx [agg_key ]] += sq 53 layer[self._layer_idx[layer_key]] += sq 54 any_grad = True 55 56 if not any_grad: 57 return {}, {} 58 59 # Single batched CPU-GPU sync 60 all_norms = torch.cat(agg + layer).sqrt().cpu() 61 split = len(self._agg_keys) 62 63 aggregated = {k: all_norms[i].item() for i, k in enumerate(self._agg_keys)} 64 per_layer = {k: all_norms[split + i].item() for i, k in enumerate(self._layer_keys)} 65 66 return aggregated, per_layer 67 68def get_weight_norms( named_parameters ): 69 70 norms = {} 71 72 target_keywords = ( 73 "embed", "wte", "wpe", # Embeddings 74 "attn", "q_proj", "k_proj", "v_proj", # Attention 75 "o_proj", "c_attn", "qkv", 76 "mlp", "gate_proj", "up_proj", # MLPs 77 "down_proj", "c_fc", "fc", "dense", 78 "lm_head", "score" # Heads 79 ) 80 81 for name, param in named_parameters : 82 83 if "bias" in name: 84 continue 85 86 # Specifically monitor the embedding and head (output) 87 if any(keyword in name for keyword in target_keywords) : 88 89 # Replace dots with underscores for cleaner Aim UI paths 90 safe_name = name.replace(".", "_") 91 norms[f"weights/norm_{safe_name}"] = torch.norm(param, p='fro').item() 92 93 return norms 94 95@dataclass 96class WindowMetrics: 97 mean_loss: float 98 grad_norm: float 99 lr: float 100 tok_per_s: float 101 ms_per_step: float 102 skipped_tokens: int 103 mean_block_norms: dict 104 per_layer_norms: dict 105 per_block_norms: dict 106 extra_loss_terms: dict = field( default_factory=dict ) 107 108@dataclass 109class Accumulator : 110 111 token_weighted_window_loss : float = 0.0 112 windowed_tokens : int = 0 113 windowed_skipped_tokens : int = 0 114 windowed_grad_norm : float = 0.0 115 windowed_block_norms : dict = field( default_factory=dict ) 116 windowed_steps : int = 0 117 windowed_all_steps : int = 0 118 last_per_layer_norms : dict = field( default_factory=dict ) 119 last_per_block_norms : dict = field( default_factory=dict ) 120 windowed_extra_terms : dict = field( default_factory=dict ) 121 122 def update(self, loss_sum_weighted, step_tokens, grad_norm, skipped, block_norms, extra_terms=None ): 123 124 """Absorbs the raw output from a training step.""" 125 self.windowed_all_steps += 1 126 127 if skipped: 128 self.windowed_skipped_tokens += step_tokens 129 return 130 131 self.token_weighted_window_loss += loss_sum_weighted 132 self.windowed_tokens += step_tokens 133 self.windowed_grad_norm += grad_norm 134 self.windowed_steps += 1 135 136 agg_norms, per_layer_norms = block_norms 137 138 for k, v in agg_norms.items(): 139 self.windowed_block_norms[k] = self.windowed_block_norms.get(k, 0.0) + v 140 141 self.last_per_layer_norms = per_layer_norms 142 self.last_per_block_norms = agg_norms 143 144 for k, v in (extra_terms or {}).items(): 145 self.windowed_extra_terms[k] = self.windowed_extra_terms.get(k, 0.0) + v 146 147 def to_metrics(self, elapsed_time: float, lr_now: float) -> WindowMetrics: 148 """Calculates averages and returns the structured dataclass.""" 149 n = max(self.windowed_steps, 1) 150 151 return WindowMetrics( 152 mean_loss=self.token_weighted_window_loss / max(self.windowed_tokens, 1), 153 grad_norm=self.windowed_grad_norm / n, 154 lr=lr_now, 155 tok_per_s=self.windowed_tokens / max(elapsed_time, 1e-6), 156 ms_per_step=elapsed_time * 1000.0 / max(self.windowed_all_steps, 1), 157 skipped_tokens=self.windowed_skipped_tokens, 158 mean_block_norms={k: v / n for k, v in self.windowed_block_norms.items()}, 159 per_layer_norms=self.last_per_layer_norms, 160 per_block_norms=self.last_per_block_norms, 161 extra_loss_terms={k: v / n for k, v in self.windowed_extra_terms.items()}, 162 )
def
block_key(name: str) -> tuple[int | None, str]:
9def block_key(name: str) -> tuple[ int | None, str ]: 10 parts = name.split(".") 11 for i, part in enumerate(parts): 12 if part in _LAYER_CONTAINERS and i + 2 < len(parts) and parts[i + 1].isdigit(): 13 block = parts[i + 2] 14 component = parts[0] 15 return int(parts[i + 1]), f"{component}/{block}" 16 # No layer index — use the direct parent module as the key 17 return None, parts[-2] if len(parts) > 1 else parts[0]
class
GradNormTracker:
19class GradNormTracker: 20 """Caches the parameter→block-key mapping once at init, then just reads grads each call.""" 21 22 def __init__(self, model, device): 23 self.device = device 24 # Pre-compute (param_ref, agg_key, layer_key) for every parameter once. 25 self._entries: list[tuple[torch.nn.Parameter, str, str]] = [] 26 for name, p in model.named_parameters(): 27 layer_idx, block = block_key(name) 28 agg_key = f"grad_norm/{block}" 29 layer_key = (f"grad_norm/layer_{layer_idx:03d}/{block}" 30 if layer_idx is not None 31 else f"grad_norm/model/{block}") 32 self._entries.append((p, agg_key, layer_key)) 33 34 # Pre-sort the unique key lists so the batched sync is always in a stable order. 35 seen_agg, seen_layer = {}, {} 36 for _, agg_key, layer_key in self._entries: 37 seen_agg[agg_key] = True 38 seen_layer[layer_key] = True 39 self._agg_keys = sorted(seen_agg) 40 self._layer_keys = sorted(seen_layer) 41 self._agg_idx = {k: i for i, k in enumerate(self._agg_keys)} 42 self._layer_idx = {k: i for i, k in enumerate(self._layer_keys)} 43 44 def compute(self) -> tuple[dict, dict]: 45 agg = [torch.zeros(1, device=self.device) for _ in self._agg_keys] 46 layer = [torch.zeros(1, device=self.device) for _ in self._layer_keys] 47 48 any_grad = False 49 for p, agg_key, layer_key in self._entries: 50 if p.grad is None: 51 continue 52 sq = p.grad.detach().float().pow(2).sum() 53 agg [self._agg_idx [agg_key ]] += sq 54 layer[self._layer_idx[layer_key]] += sq 55 any_grad = True 56 57 if not any_grad: 58 return {}, {} 59 60 # Single batched CPU-GPU sync 61 all_norms = torch.cat(agg + layer).sqrt().cpu() 62 split = len(self._agg_keys) 63 64 aggregated = {k: all_norms[i].item() for i, k in enumerate(self._agg_keys)} 65 per_layer = {k: all_norms[split + i].item() for i, k in enumerate(self._layer_keys)} 66 67 return aggregated, per_layer
Caches the parameter→block-key mapping once at init, then just reads grads each call.
GradNormTracker(model, device)
22 def __init__(self, model, device): 23 self.device = device 24 # Pre-compute (param_ref, agg_key, layer_key) for every parameter once. 25 self._entries: list[tuple[torch.nn.Parameter, str, str]] = [] 26 for name, p in model.named_parameters(): 27 layer_idx, block = block_key(name) 28 agg_key = f"grad_norm/{block}" 29 layer_key = (f"grad_norm/layer_{layer_idx:03d}/{block}" 30 if layer_idx is not None 31 else f"grad_norm/model/{block}") 32 self._entries.append((p, agg_key, layer_key)) 33 34 # Pre-sort the unique key lists so the batched sync is always in a stable order. 35 seen_agg, seen_layer = {}, {} 36 for _, agg_key, layer_key in self._entries: 37 seen_agg[agg_key] = True 38 seen_layer[layer_key] = True 39 self._agg_keys = sorted(seen_agg) 40 self._layer_keys = sorted(seen_layer) 41 self._agg_idx = {k: i for i, k in enumerate(self._agg_keys)} 42 self._layer_idx = {k: i for i, k in enumerate(self._layer_keys)}
def
compute(self) -> tuple[dict, dict]:
44 def compute(self) -> tuple[dict, dict]: 45 agg = [torch.zeros(1, device=self.device) for _ in self._agg_keys] 46 layer = [torch.zeros(1, device=self.device) for _ in self._layer_keys] 47 48 any_grad = False 49 for p, agg_key, layer_key in self._entries: 50 if p.grad is None: 51 continue 52 sq = p.grad.detach().float().pow(2).sum() 53 agg [self._agg_idx [agg_key ]] += sq 54 layer[self._layer_idx[layer_key]] += sq 55 any_grad = True 56 57 if not any_grad: 58 return {}, {} 59 60 # Single batched CPU-GPU sync 61 all_norms = torch.cat(agg + layer).sqrt().cpu() 62 split = len(self._agg_keys) 63 64 aggregated = {k: all_norms[i].item() for i, k in enumerate(self._agg_keys)} 65 per_layer = {k: all_norms[split + i].item() for i, k in enumerate(self._layer_keys)} 66 67 return aggregated, per_layer
def
get_weight_norms(named_parameters):
69def get_weight_norms( named_parameters ): 70 71 norms = {} 72 73 target_keywords = ( 74 "embed", "wte", "wpe", # Embeddings 75 "attn", "q_proj", "k_proj", "v_proj", # Attention 76 "o_proj", "c_attn", "qkv", 77 "mlp", "gate_proj", "up_proj", # MLPs 78 "down_proj", "c_fc", "fc", "dense", 79 "lm_head", "score" # Heads 80 ) 81 82 for name, param in named_parameters : 83 84 if "bias" in name: 85 continue 86 87 # Specifically monitor the embedding and head (output) 88 if any(keyword in name for keyword in target_keywords) : 89 90 # Replace dots with underscores for cleaner Aim UI paths 91 safe_name = name.replace(".", "_") 92 norms[f"weights/norm_{safe_name}"] = torch.norm(param, p='fro').item() 93 94 return norms
@dataclass
class
WindowMetrics:
96@dataclass 97class WindowMetrics: 98 mean_loss: float 99 grad_norm: float 100 lr: float 101 tok_per_s: float 102 ms_per_step: float 103 skipped_tokens: int 104 mean_block_norms: dict 105 per_layer_norms: dict 106 per_block_norms: dict 107 extra_loss_terms: dict = field( default_factory=dict )
@dataclass
class
Accumulator:
109@dataclass 110class Accumulator : 111 112 token_weighted_window_loss : float = 0.0 113 windowed_tokens : int = 0 114 windowed_skipped_tokens : int = 0 115 windowed_grad_norm : float = 0.0 116 windowed_block_norms : dict = field( default_factory=dict ) 117 windowed_steps : int = 0 118 windowed_all_steps : int = 0 119 last_per_layer_norms : dict = field( default_factory=dict ) 120 last_per_block_norms : dict = field( default_factory=dict ) 121 windowed_extra_terms : dict = field( default_factory=dict ) 122 123 def update(self, loss_sum_weighted, step_tokens, grad_norm, skipped, block_norms, extra_terms=None ): 124 125 """Absorbs the raw output from a training step.""" 126 self.windowed_all_steps += 1 127 128 if skipped: 129 self.windowed_skipped_tokens += step_tokens 130 return 131 132 self.token_weighted_window_loss += loss_sum_weighted 133 self.windowed_tokens += step_tokens 134 self.windowed_grad_norm += grad_norm 135 self.windowed_steps += 1 136 137 agg_norms, per_layer_norms = block_norms 138 139 for k, v in agg_norms.items(): 140 self.windowed_block_norms[k] = self.windowed_block_norms.get(k, 0.0) + v 141 142 self.last_per_layer_norms = per_layer_norms 143 self.last_per_block_norms = agg_norms 144 145 for k, v in (extra_terms or {}).items(): 146 self.windowed_extra_terms[k] = self.windowed_extra_terms.get(k, 0.0) + v 147 148 def to_metrics(self, elapsed_time: float, lr_now: float) -> WindowMetrics: 149 """Calculates averages and returns the structured dataclass.""" 150 n = max(self.windowed_steps, 1) 151 152 return WindowMetrics( 153 mean_loss=self.token_weighted_window_loss / max(self.windowed_tokens, 1), 154 grad_norm=self.windowed_grad_norm / n, 155 lr=lr_now, 156 tok_per_s=self.windowed_tokens / max(elapsed_time, 1e-6), 157 ms_per_step=elapsed_time * 1000.0 / max(self.windowed_all_steps, 1), 158 skipped_tokens=self.windowed_skipped_tokens, 159 mean_block_norms={k: v / n for k, v in self.windowed_block_norms.items()}, 160 per_layer_norms=self.last_per_layer_norms, 161 per_block_norms=self.last_per_block_norms, 162 extra_loss_terms={k: v / n for k, v in self.windowed_extra_terms.items()}, 163 )
Accumulator( token_weighted_window_loss: float = 0.0, windowed_tokens: int = 0, windowed_skipped_tokens: int = 0, windowed_grad_norm: float = 0.0, windowed_block_norms: dict = <factory>, windowed_steps: int = 0, windowed_all_steps: int = 0, last_per_layer_norms: dict = <factory>, last_per_block_norms: dict = <factory>, windowed_extra_terms: dict = <factory>)
def
update( self, loss_sum_weighted, step_tokens, grad_norm, skipped, block_norms, extra_terms=None):
123 def update(self, loss_sum_weighted, step_tokens, grad_norm, skipped, block_norms, extra_terms=None ): 124 125 """Absorbs the raw output from a training step.""" 126 self.windowed_all_steps += 1 127 128 if skipped: 129 self.windowed_skipped_tokens += step_tokens 130 return 131 132 self.token_weighted_window_loss += loss_sum_weighted 133 self.windowed_tokens += step_tokens 134 self.windowed_grad_norm += grad_norm 135 self.windowed_steps += 1 136 137 agg_norms, per_layer_norms = block_norms 138 139 for k, v in agg_norms.items(): 140 self.windowed_block_norms[k] = self.windowed_block_norms.get(k, 0.0) + v 141 142 self.last_per_layer_norms = per_layer_norms 143 self.last_per_block_norms = agg_norms 144 145 for k, v in (extra_terms or {}).items(): 146 self.windowed_extra_terms[k] = self.windowed_extra_terms.get(k, 0.0) + v
Absorbs the raw output from a training step.
148 def to_metrics(self, elapsed_time: float, lr_now: float) -> WindowMetrics: 149 """Calculates averages and returns the structured dataclass.""" 150 n = max(self.windowed_steps, 1) 151 152 return WindowMetrics( 153 mean_loss=self.token_weighted_window_loss / max(self.windowed_tokens, 1), 154 grad_norm=self.windowed_grad_norm / n, 155 lr=lr_now, 156 tok_per_s=self.windowed_tokens / max(elapsed_time, 1e-6), 157 ms_per_step=elapsed_time * 1000.0 / max(self.windowed_all_steps, 1), 158 skipped_tokens=self.windowed_skipped_tokens, 159 mean_block_norms={k: v / n for k, v in self.windowed_block_norms.items()}, 160 per_layer_norms=self.last_per_layer_norms, 161 per_block_norms=self.last_per_block_norms, 162 extra_loss_terms={k: v / n for k, v in self.windowed_extra_terms.items()}, 163 )
Calculates averages and returns the structured dataclass.