evalmetry 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
evalmetry/reducers.py ADDED
@@ -0,0 +1,578 @@
1
+ """Reducers: turn the tensors a hook catches into a handful of numbers.
2
+
3
+ A reducer is the only place a raw activation is allowed to live.
4
+ Hooks hand it a tensor, it keeps just what it needs, and it drops the reference immediately - vocab-sized logits and full hidden states cannot be written out at evaluation scale, so the shrinking has to happen inside the forward pass.
5
+
6
+ One signal = one reducer.
7
+ Adding a new signal means adding a class here; no hook code and no lm-eval glue has to change.
8
+
9
+ Lifecycle, driven by `recorder.py`:
10
+
11
+ update(index, tensor, ctx) once per captured tensor (per layer, per hook) end_forward(ctx) every layer of one forward pass has arrived; do the maths now and let the tensors go finalize() the document is done; hand back rows and reset
12
+
13
+ The `update` / `finalize` pair is the reducer interface; `end_forward` exists because a generate task runs one forward pass per decoding step, and each step has to be reduced on the spot rather than piled up until the document ends - the whole point is never to hold vocab-sized tensors.
14
+
15
+ Every reducer declares three things about its output, which `report.py` reads to decide what may share an axis:
16
+
17
+ comparable_across_vocab may models with different vocabularies be overlaid? layer_domain "residual" (0..L, column `layer`) or "block" (0..L-1, column `block`) version bump whenever the *definition* of the number changes, even if the column names do not
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ from dataclasses import dataclass, field
23
+ from typing import Any, Callable, Sequence
24
+
25
+ import torch
26
+
27
+
28
+ # : Positions are decoded in chunks so the transient (chunk x layers x vocab) : logit tensor stays bounded. 33 layers x 150k vocab x fp32 is ~20 MB per : position, which is fine one at a time and not fine for a long scored span.
29
+ LENS_POSITION_CHUNK = 4
30
+
31
+
32
+ @dataclass
33
+ class ForwardContext:
34
+ """Everything the reducers need to know about the forward pass in flight.
35
+
36
+ One instance describes one forward pass of one (document, choice).
37
+ For a loglikelihood request that is the whole document; for a generate request there is one context per decoding step.
38
+
39
+ Attributes:
40
+ task_name: lm-eval task name.
41
+ doc_id: lm-eval document id.
42
+ choice_idx: which choice of a multiple-choice document, else 0.
43
+ steps: `step` axis value for each position carried by this forward. loglikelihood: 0..contlen-1, the positions lm-eval actually scores. generate: a single decoding step, e.g. [7].
44
+ positions: absolute position in the input sequence for each step.
45
+ seq_indices: index along the tensor's sequence axis for each step.
46
+ Usually the same as `positions`, but a generate decode step feeds a single-token tensor, where the only valid index is 0.
47
+ batch_row: which row of the batched tensors belongs to this document.
48
+ target_token_ids: gold continuation token at each step, or None for a generate task where there is no gold token to rank.
49
+ step_token_ids: token to report in the `steps` table.
50
+ Identical to `target_token_ids` for loglikelihood; for generate it is the token that was actually emitted, which sampling can make differ from the last layer's top-1.
51
+ n_residual: number of residual entries, L+1.
52
+ n_blocks: number of transformer blocks, L.
53
+ task_kind: "loglikelihood" or "generate".
54
+ Reducers that store data sparsely along the step axis (attention weights) use it to decide which steps to keep.
55
+ shared: scratch space reducers may use to pass intermediate values to each other within one forward pass.
56
+ Cleared by the recorder afterwards, so nothing survives into the next pass.
57
+
58
+ Example:
59
+ >>> ctx = ForwardContext("xnli_ko", 7, 1, steps=[0], positions=[41],
60
+ ... target_token_ids=[9891], n_residual=33, n_blocks=32)
61
+ >>> ctx.key
62
+ ('xnli_ko', 7, 1)
63
+ >>> ctx.n_positions
64
+ 1
65
+ >>> ctx.seq_indices
66
+ [41]
67
+ """
68
+
69
+ task_name: str
70
+ doc_id: int
71
+ choice_idx: int
72
+ steps: list[int]
73
+ positions: list[int]
74
+ target_token_ids: list[int] | None
75
+ n_residual: int
76
+ n_blocks: int
77
+ task_kind: str = "loglikelihood"
78
+ batch_row: int = 0
79
+ seq_indices: list[int] | None = None
80
+ step_token_ids: list[int] | None = None
81
+ shared: dict[str, Any] = field(default_factory=dict)
82
+ # Actual unpadded model input length, not the length of the scored span.
83
+ # Module statistics need this to exclude padding across every input position.
84
+ input_length: int | None = None
85
+ # Tokens removed from the left before the actual model input; None means unknown.
86
+ input_offset: int | None = None
87
+
88
+ def __post_init__(self) -> None:
89
+ # By default a step reads the tensor at its own absolute position; only incremental decoding needs a different mapping.
90
+ if self.seq_indices is None:
91
+ self.seq_indices = list(self.positions)
92
+ if self.step_token_ids is None:
93
+ self.step_token_ids = self.target_token_ids
94
+
95
+ @property
96
+ def key(self) -> tuple[str, int, int]:
97
+ """The (task_name, doc_id, choice_idx) triple every row is keyed by."""
98
+ return (self.task_name, self.doc_id, self.choice_idx)
99
+
100
+ @property
101
+ def n_positions(self) -> int:
102
+ return len(self.steps)
103
+
104
+ def axis_columns(self, step_index: int) -> dict[str, Any]:
105
+ """The four shared axis columns for one position of this forward pass.
106
+
107
+ Example:
108
+ >>> ctx = ForwardContext("t", 3, 0, [0, 1], [10, 11], None, 33, 32)
109
+ >>> ctx.axis_columns(1)
110
+ {'task_name': 't', 'doc_id': 3, 'choice_idx': 0, 'step': 1}
111
+ """
112
+ return {
113
+ "task_name": self.task_name,
114
+ "doc_id": self.doc_id,
115
+ "choice_idx": self.choice_idx,
116
+ "step": self.steps[step_index],
117
+ }
118
+
119
+
120
+ class Reducer:
121
+ """Base class.
122
+ Subclasses override `update`, `end_forward` and `finalize`.
123
+
124
+ The defaults are no-ops so a reducer only has to implement the callbacks it actually cares about: a residual-stream reducer ignores the attention hooks, and an attention reducer ignores the residual ones.
125
+ """
126
+
127
+ #: Short identifier written into the manifest.
128
+ name: str = "reducer"
129
+ #: Bump when the meaning of the numbers changes.
130
+ version: int = 1
131
+ # : Which table in `storage.TABLES` the rows go to, or None for a reducer : that only writes tensor files.
132
+ table: str | None = None
133
+ #: "residual" (0..L) or "block" (0..L-1), or None if it has no layer axis.
134
+ layer_domain: str | None = None
135
+ #: Whether its values may be plotted next to a model with another vocab.
136
+ comparable_across_vocab: bool = True
137
+ #: Which hook stream feeds it: "residual", "attention" or "value".
138
+ source: str = "residual"
139
+
140
+ def config(self) -> dict[str, Any]:
141
+ """Settings recorded in the manifest, so an old run stays interpretable."""
142
+ return {}
143
+
144
+ def descriptor(self) -> dict[str, Any]:
145
+ """The manifest entry for this reducer.
146
+
147
+ Example:
148
+ >>> SimilarityReducer(n_residual=3).descriptor()["name"]
149
+ 'layer_similarity'
150
+ """
151
+ return {
152
+ "name": self.name,
153
+ "version": self.version,
154
+ "table": self.table,
155
+ "layer_domain": self.layer_domain,
156
+ "comparable_across_vocab": self.comparable_across_vocab,
157
+ "config": self.config(),
158
+ }
159
+
160
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
161
+ """One captured tensor.
162
+ `index` is a layer index or a block index."""
163
+
164
+ def end_forward(self, ctx: ForwardContext) -> None:
165
+ """Every tensor of this forward pass has arrived: reduce and release."""
166
+
167
+ def finalize(self) -> list[dict[str, Any]]:
168
+ """The document is finished: return its rows and clear internal state."""
169
+ return []
170
+
171
+ def pop_tensors(self) -> dict[str, torch.Tensor]:
172
+ """Tensors to be written to safetensors for this document, if any.
173
+
174
+ Only the opt-in dumpers return anything here; the recorder hands the result to `storage.TensorStore`.
175
+ """
176
+ return {}
177
+
178
+
179
+ # --------------------------------------------------------------------------
180
+ # Always-on signals
181
+ # --------------------------------------------------------------------------
182
+
183
+
184
+ class LogitLensReducer(Reducer):
185
+ """Logit lens: read every layer's residual stream through the model's own head.
186
+
187
+ For each (step, layer) it records the top-1 token and, for loglikelihood tasks, where the gold token sits in that layer's ranking.
188
+
189
+ Why one reducer and not two Plan section 4 lists "logit lens" and "gold token rank" as separate signals.
190
+ They are computed together here because both are functions of the *same* layer-wise logit tensor, and that tensor is the one thing we are not allowed to keep alive (~20 MB per position at 33 layers x 150k vocab).
191
+ Handing it to a second reducer would mean either holding it longer or paying for a second unembedding GEMM.
192
+ The two signals therefore share this reducer's `version`.
193
+
194
+ Why k is fixed at 1 With k=1 there is exactly one row per (doc_id, choice_idx, step, layer), so the table needs no extra k axis and lines up with every other signal.
195
+ "Where is the gold token?" is answered by `target_rank`, not by a top-k list.
196
+ Raising k adds an axis, so it would need a `version` bump.
197
+
198
+ Args:
199
+ decode_stack: `(n_layers, n_pos, d) -> (n_layers, n_pos, vocab)`.
200
+ Supplied by the adapter, because the path from a hidden state to logits (final norm, head bias, logit softcapping, output scaling) differs per architecture and must not be reimplemented here.
201
+ tokenizer: used only to turn ids into strings and to flag specials.
202
+ vocab_size: denominator of `target_percentile`.
203
+
204
+ Example of one produced row:
205
+ {"task_name": "xnli_ko", "doc_id": 7, "choice_idx": 1, "step": 0, "layer": 20, "lens_token_id": 9891, "lens_prob": 0.31, "lens_token": "Ġyes", "lens_is_special": False, "target_rank": 2, "target_percentile": 1.3e-05}
206
+ """
207
+
208
+ name = "logit_lens"
209
+ version = 1
210
+ table = "signals"
211
+ layer_domain = "residual"
212
+ comparable_across_vocab = False # token ids and probabilities are vocab-bound
213
+ source = "residual"
214
+
215
+ def __init__(
216
+ self,
217
+ decode_stack: Callable[[torch.Tensor], torch.Tensor],
218
+ tokenizer: Any,
219
+ vocab_size: int,
220
+ position_chunk: int = LENS_POSITION_CHUNK,
221
+ ) -> None:
222
+ self._decode_stack = decode_stack
223
+ self._tokenizer = tokenizer
224
+ self._vocab_size = int(vocab_size)
225
+ self._position_chunk = position_chunk
226
+ self._special_ids = set(getattr(tokenizer, "all_special_ids", []) or [])
227
+ self._token_text: dict[int, str] = {} # id -> string, cached for the whole run
228
+ self._layers: dict[int, torch.Tensor] = {}
229
+ self._rows: list[dict[str, Any]] = []
230
+
231
+ def config(self) -> dict[str, Any]:
232
+ return {"top_k": 1, "vocab_size": self._vocab_size, "position_chunk": self._position_chunk}
233
+
234
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
235
+ # `tensor` is already sliced down to the scored positions, (n_pos, d), so what we hold on to between here and end_forward is small.
236
+ self._layers[index] = tensor
237
+
238
+ def end_forward(self, ctx: ForwardContext) -> None:
239
+ if not self._layers:
240
+ return
241
+ stack = torch.stack([self._layers[i] for i in range(ctx.n_residual)], dim=0)
242
+ self._layers.clear()
243
+
244
+ for start in range(0, ctx.n_positions, self._position_chunk):
245
+ stop = min(start + self._position_chunk, ctx.n_positions)
246
+ chunk = stack[:, start:stop, :] # (n_layers, chunk, d)
247
+ # One GEMM for the whole layer stack.
248
+ # Looping layer by layer costs noticeably more on generate tasks, where this runs every step.
249
+ logits = self._decode_stack(chunk).float() # (n_layers, chunk, vocab)
250
+ logprobs = torch.log_softmax(logits, dim=-1)
251
+ top_prob, top_id = logprobs.max(dim=-1)
252
+ top_prob = top_prob.exp()
253
+
254
+ gold_rank = None
255
+ if ctx.target_token_ids is not None:
256
+ gold = torch.tensor(
257
+ ctx.target_token_ids[start:stop], device=logits.device, dtype=torch.long
258
+ )
259
+ gold_logit = logits.gather(
260
+ -1, gold.view(1, -1, 1).expand(logits.shape[0], -1, 1)
261
+ )
262
+ # 0-based rank: how many tokens beat the gold token.
263
+ # Ties count as not beating it, so the rank is the optimistic one.
264
+ gold_rank = (logits > gold_logit).sum(dim=-1)
265
+
266
+ del logits, logprobs
267
+ self._emit(ctx, start, stop, top_id.cpu(), top_prob.cpu(),
268
+ None if gold_rank is None else gold_rank.cpu())
269
+
270
+ def _emit(
271
+ self,
272
+ ctx: ForwardContext,
273
+ start: int,
274
+ stop: int,
275
+ top_id: torch.Tensor,
276
+ top_prob: torch.Tensor,
277
+ gold_rank: torch.Tensor | None,
278
+ ) -> None:
279
+ """Turn the reduced tensors of one position chunk into rows."""
280
+ ids = top_id.tolist()
281
+ probs = top_prob.tolist()
282
+ ranks = gold_rank.tolist() if gold_rank is not None else None
283
+ for local_pos, pos in enumerate(range(start, stop)):
284
+ axis = ctx.axis_columns(pos)
285
+ for layer in range(ctx.n_residual):
286
+ token_id = int(ids[layer][local_pos])
287
+ row = dict(axis)
288
+ row.update(
289
+ layer=layer,
290
+ lens_token_id=token_id,
291
+ lens_prob=float(probs[layer][local_pos]),
292
+ lens_token=self._text_for(token_id),
293
+ lens_is_special=token_id in self._special_ids,
294
+ target_rank=None,
295
+ target_percentile=None,
296
+ )
297
+ if ranks is not None:
298
+ rank = int(ranks[layer][local_pos])
299
+ row["target_rank"] = rank
300
+ row["target_percentile"] = rank / self._vocab_size
301
+ self._rows.append(row)
302
+
303
+ def _text_for(self, token_id: int) -> str:
304
+ """Tokenizer string for a token id, stored verbatim.
305
+
306
+ Uses `convert_ids_to_tokens`, i.e. the raw vocabulary piece, so leading space markers ("Ġ") and partial byte pieces survive untouched.
307
+ Cleaning them up is left to whoever draws the figure.
308
+ """
309
+ cached = self._token_text.get(token_id)
310
+ if cached is None:
311
+ try:
312
+ cached = self._tokenizer.convert_ids_to_tokens(token_id)
313
+ except Exception: # tokenizers without the fast-tokenizer API
314
+ cached = self._tokenizer.decode([token_id])
315
+ cached = "" if cached is None else str(cached)
316
+ self._token_text[token_id] = cached
317
+ return cached
318
+
319
+ def finalize(self) -> list[dict[str, Any]]:
320
+ rows, self._rows = self._rows, []
321
+ self._layers.clear()
322
+ return rows
323
+
324
+
325
+ class SimilarityReducer(Reducer):
326
+ """Cosine similarity between every pair of residual layers, per position.
327
+
328
+ Stores the upper triangle including the diagonal, since the matrix is symmetric: 561 values for a 32-block model, about 2.2 KB per step in fp32.
329
+
330
+ No centering is applied.
331
+ The residual stream is only ever added to, so all tokens share a large common component and the cosines sit close to 1; read the *relative* structure between layer pairs, not the absolute value.
332
+ Subtracting a corpus mean needs a second pass and is left as a TODO, which is exactly the kind of change that would require a `version` bump.
333
+
334
+ Example of one produced row:
335
+ {"task_name": "xnli_ko", "doc_id": 7, "choice_idx": 1, "step": 0, "layer_i": 3, "layer_j": 17, "cos": 0.94}
336
+ """
337
+
338
+ name = "layer_similarity"
339
+ version = 1
340
+ table = "similarity"
341
+ layer_domain = "residual"
342
+ comparable_across_vocab = True # a cosine does not depend on vocabulary size
343
+ source = "residual"
344
+
345
+ def __init__(self, n_residual: int) -> None:
346
+ # The (i, j) index pairs are fixed for a model, so build them once.
347
+ rows, cols = torch.triu_indices(n_residual, n_residual, offset=0).tolist()
348
+ self._pair_i = rows
349
+ self._pair_j = cols
350
+ self._layers: dict[int, torch.Tensor] = {}
351
+ self._rows: list[dict[str, Any]] = []
352
+
353
+ def config(self) -> dict[str, Any]:
354
+ return {"centering": False, "store": "upper_triangle_with_diagonal"}
355
+
356
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
357
+ self._layers[index] = tensor
358
+
359
+ def end_forward(self, ctx: ForwardContext) -> None:
360
+ if not self._layers:
361
+ return
362
+ stack = torch.stack([self._layers[i] for i in range(ctx.n_residual)], dim=0).float()
363
+ self._layers.clear()
364
+ # (n_layers, n_pos, d) -> unit vectors -> per position Gram matrix.
365
+ normed = torch.nn.functional.normalize(stack, dim=-1)
366
+ for pos in range(ctx.n_positions):
367
+ vectors = normed[:, pos, :] # (n_layers, d)
368
+ matrix = (vectors @ vectors.T).cpu() # (n_layers, n_layers)
369
+ values = matrix[self._pair_i, self._pair_j].tolist()
370
+ axis = ctx.axis_columns(pos)
371
+ for layer_i, layer_j, cos in zip(self._pair_i, self._pair_j, values):
372
+ row = dict(axis)
373
+ row.update(layer_i=layer_i, layer_j=layer_j, cos=float(cos))
374
+ self._rows.append(row)
375
+
376
+ def finalize(self) -> list[dict[str, Any]]:
377
+ rows, self._rows = self._rows, []
378
+ self._layers.clear()
379
+ return rows
380
+
381
+
382
+ # --------------------------------------------------------------------------
383
+ # Opt-in signals (--save-attention, --save-hidden)
384
+ # --------------------------------------------------------------------------
385
+
386
+
387
+ class ValueNormReducer(Reducer):
388
+ """L2 norm of each attention head's value vector at the current position.
389
+
390
+ Cheap enough to keep on every step: one scalar per head, so 32 blocks x 32 heads in fp32 is 4 KB per step.
391
+ That is why `attn_norm` is dense along the step axis while the attention weight tensors next to it are sparse.
392
+
393
+ It has to be captured during the forward pass - the value vectors are gone afterwards, and nothing in the saved attention weights lets you recover them.
394
+
395
+ GQA note: several query heads share one key/value head.
396
+ We expand the per-kv-head norm back out to query heads so this table joins cleanly with the attention weights on (block, head); the value therefore repeats within a group.
397
+
398
+ Example of one produced row:
399
+ {"task_name": "xnli_ko", "doc_id": 7, "choice_idx": 1, "step": 0, "block": 12, "head": 5, "value_norm": 1.84}
400
+ """
401
+
402
+ name = "value_norm"
403
+ version = 1
404
+ table = "attn_norm"
405
+ layer_domain = "block"
406
+ comparable_across_vocab = True
407
+ source = "value"
408
+
409
+ def __init__(
410
+ self,
411
+ n_heads: int,
412
+ n_kv_heads: int | None,
413
+ head_dim: int | None,
414
+ block_shapes: dict[int, tuple[int, int, int]] | None = None,
415
+ ) -> None:
416
+ # A model whose blocks differ in attention shape passes `block_shapes`, block -> (n_heads, n_kv_heads, head_dim).
417
+ # Gemma 4 does: its sliding blocks use 256-wide heads and its full-attention blocks 512-wide ones, so one reshape for every block would split a 1024-column value row into the wrong heads.
418
+ self._n_heads = n_heads
419
+ self._n_kv_heads = n_kv_heads
420
+ self._head_dim = head_dim
421
+ self._block_shapes = dict(block_shapes or {})
422
+ for heads, kv_heads, _ in set(self._block_shapes.values()) or {(n_heads, n_kv_heads, head_dim)}:
423
+ if not kv_heads or heads % kv_heads != 0:
424
+ raise ValueError(
425
+ f"n_heads ({heads}) is not a multiple of n_kv_heads ({kv_heads}); "
426
+ "cannot map query heads onto key/value heads"
427
+ )
428
+ self._rows: list[dict[str, Any]] = []
429
+
430
+ def config(self) -> dict[str, Any]:
431
+ entry: dict[str, Any] = {
432
+ "n_heads": self._n_heads,
433
+ "n_kv_heads": self._n_kv_heads,
434
+ "head_dim": self._head_dim,
435
+ "gqa_expansion": "repeat_interleave to query heads",
436
+ }
437
+ if self._block_shapes:
438
+ entry["shapes_by_block"] = {
439
+ str(block): list(shape) for block, shape in sorted(self._block_shapes.items())}
440
+ return entry
441
+
442
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
443
+ """`tensor` is the v_proj output at the scored positions, (n_pos, n_kv*head_dim)."""
444
+ if self._block_shapes:
445
+ if index not in self._block_shapes:
446
+ raise KeyError(f"block {index} has a value projection but no recorded attention shape")
447
+ n_heads, n_kv_heads, head_dim = self._block_shapes[index]
448
+ else:
449
+ n_heads, n_kv_heads, head_dim = self._n_heads, self._n_kv_heads, self._head_dim
450
+ values = tensor.float().view(tensor.shape[0], n_kv_heads, head_dim)
451
+ norms = values.norm(dim=-1) # (n_pos, n_kv_heads)
452
+ norms = norms.repeat_interleave(n_heads // n_kv_heads, dim=-1) # (n_pos, n_heads)
453
+ norms = norms.cpu().tolist()
454
+ for pos in range(ctx.n_positions):
455
+ axis = ctx.axis_columns(pos)
456
+ for head in range(n_heads):
457
+ row = dict(axis)
458
+ row.update(block=index, head=head, value_norm=float(norms[pos][head]))
459
+ self._rows.append(row)
460
+
461
+ def finalize(self) -> list[dict[str, Any]]:
462
+ rows, self._rows = self._rows, []
463
+ return rows
464
+
465
+
466
+ class AttentionWeightReducer(Reducer):
467
+ """Dumps the last position's attention weights to safetensors.
468
+
469
+ Storage, not reduction: the weights go out as `(blocks, heads, seq)` per saved step.
470
+ A full (seq x seq) map would be ~537 MB per document at seq=512; keeping only the last query position brings that to ~1 MB.
471
+
472
+ The step axis keeps its usual meaning, but only a subset of steps is written:
473
+
474
+ * loglikelihood - every step of the scored span (usually one, at most a few)
475
+ * generate - only `step 0`, the last prefill position, because a 100-step generation would otherwise cost ~100 MB per document, orders of magnitude more than every other signal
476
+
477
+ Attention sinks are left in.
478
+ The report labels this axis "attention weight", never "contribution" or "reference".
479
+ """
480
+
481
+ name = "attention_weights"
482
+ version = 1
483
+ table = None # writes tensor files, not parquet rows
484
+ layer_domain = "block"
485
+ comparable_across_vocab = True
486
+ source = "attention"
487
+
488
+ def __init__(self) -> None:
489
+ self._per_step: dict[int, dict[int, torch.Tensor]] = {}
490
+
491
+ def config(self) -> dict[str, Any]:
492
+ return {
493
+ "position": "last query position",
494
+ "steps_saved": {"loglikelihood": "all scored steps", "generate": "step 0 only"},
495
+ "attn_implementation": "eager",
496
+ }
497
+
498
+ @staticmethod
499
+ def keeps_step(ctx: ForwardContext, step: int) -> bool:
500
+ """Whether this step's weights are written out.
501
+
502
+ Example:
503
+ >>> ctx = ForwardContext("t", 0, 0, [3], [12], None, 33, 32, task_kind="generate")
504
+ >>> AttentionWeightReducer.keeps_step(ctx, 3)
505
+ False
506
+ >>> AttentionWeightReducer.keeps_step(ctx, 0)
507
+ True
508
+ """
509
+ if ctx.task_kind == "generate":
510
+ return step == 0
511
+ return True
512
+
513
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
514
+ """`tensor` is (n_pos, n_heads, seq): the attention row of each scored position."""
515
+ for pos in range(ctx.n_positions):
516
+ step = ctx.steps[pos]
517
+ if not self.keeps_step(ctx, step):
518
+ continue
519
+ self._per_step.setdefault(step, {})[index] = tensor[pos].detach().to(torch.float16).cpu()
520
+
521
+ def pop_tensors(self) -> dict[str, torch.Tensor]:
522
+ """One entry per saved step, shaped (blocks, heads, seq).
523
+
524
+ Blocks without attention (hybrid architectures) are simply absent, so the first axis is the *recorded* blocks; `recorder.py` writes the block indices into the safetensors metadata.
525
+ """
526
+ out: dict[str, torch.Tensor] = {}
527
+ for step, by_block in sorted(self._per_step.items()):
528
+ blocks = sorted(by_block)
529
+ out[f"step_{step:04d}"] = torch.stack([by_block[b] for b in blocks], dim=0)
530
+ self._per_step.clear()
531
+ return out
532
+
533
+ def finalize(self) -> list[dict[str, Any]]:
534
+ return []
535
+
536
+
537
+ class RawHiddenReducer(Reducer):
538
+ """Dumps raw hidden states for a chosen subset of layers.
539
+
540
+ Meant for before/after comparisons around pruning and quantization, where the reduced signals are not enough and you need the vectors themselves.
541
+
542
+ This is the only signal that accepts a layer subset, because it is by far the most expensive per layer: at d=4096 in fp16 one layer costs 8 KB per position, and all layers cost ~264 KB per document.
543
+
544
+ Args:
545
+ layers: residual indices to keep, already normalised (deduplicated, ascending, range-checked) by `main.py`.
546
+ """
547
+
548
+ name = "raw_hidden"
549
+ version = 1
550
+ table = None
551
+ layer_domain = "residual"
552
+ comparable_across_vocab = True
553
+ source = "residual"
554
+
555
+ def __init__(self, layers: Sequence[int]) -> None:
556
+ self._layers = list(layers)
557
+ self._wanted = set(self._layers)
558
+ self._collected: dict[int, list[torch.Tensor]] = {}
559
+
560
+ def config(self) -> dict[str, Any]:
561
+ return {"layers": self._layers, "dtype": "fp16"}
562
+
563
+ def update(self, index: int, tensor: torch.Tensor, ctx: ForwardContext) -> None:
564
+ if index not in self._wanted:
565
+ return
566
+ self._collected.setdefault(index, []).append(tensor.detach().to(torch.float16).cpu())
567
+
568
+ def pop_tensors(self) -> dict[str, torch.Tensor]:
569
+ """One entry per kept layer, shaped (total_steps, d)."""
570
+ out = {
571
+ f"layer_{layer:03d}": torch.cat(chunks, dim=0)
572
+ for layer, chunks in sorted(self._collected.items())
573
+ }
574
+ self._collected.clear()
575
+ return out
576
+
577
+ def finalize(self) -> list[dict[str, Any]]:
578
+ return []