GitLab Repo

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)}
device
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 )
WindowMetrics( mean_loss: float, grad_norm: float, lr: float, tok_per_s: float, ms_per_step: float, skipped_tokens: int, mean_block_norms: dict, per_layer_norms: dict, per_block_norms: dict, extra_loss_terms: dict = <factory>)
mean_loss: float
grad_norm: float
lr: float
tok_per_s: float
ms_per_step: float
skipped_tokens: int
mean_block_norms: dict
per_layer_norms: dict
per_block_norms: dict
extra_loss_terms: 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>)
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
windowed_steps: int = 0
windowed_all_steps: int = 0
last_per_layer_norms: dict
last_per_block_norms: dict
windowed_extra_terms: dict
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.

def to_metrics( self, elapsed_time: float, lr_now: float) -> WindowMetrics:
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.