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/recorder.py ADDED
@@ -0,0 +1,556 @@
1
+ """Hook plumbing: catch tensors, hand them to the reducers, write the rows.
2
+
3
+ The recorder is the only object that touches PyTorch hooks for signal capture.
4
+ `debug.py` registers its own, for a different purpose and on a different schedule; the two never share a handle.
5
+ It knows which document the model is currently running, slices the batched tensors down to the positions that matter, drives the reducers, and hands their rows to storage.
6
+
7
+ One capture path, used by both task kinds
8
+ -----------------------------------------
9
+ The residual stream always comes from the decoder module's own `output_hidden_states` output, i.e. the same L+1 tensors HF returns.
10
+ A forward *pre*-hook turns the flag on and a forward hook reads the result, which means loglikelihood and generate see identical tensors.
11
+
12
+ That uniformity matters.
13
+ HF's last hidden state has the final norm already applied, while a forward hook on the last block sees the pre-norm value; mixing the two would make the logit lens apply the final norm twice, but only on the last layer, and only on one of the two task kinds.
14
+ The plot would still look reasonable.
15
+ So the path is fixed here and written into the manifest.
16
+
17
+ Note this is *not* `generate(output_hidden_states=True)`, which would keep every layer of every step alive until generation ends.
18
+ Hooking the decoder gives us one step at a time, which the reducers consume and drop immediately.
19
+
20
+ Attention weights and value vectors do come from forward hooks on the modules the adapter points at, and need `attn_implementation="eager"`.
21
+ """
22
+
23
+ from __future__ import annotations
24
+
25
+ from contextlib import contextmanager
26
+ from dataclasses import dataclass
27
+ from typing import Any, Iterator, Sequence
28
+
29
+ import torch
30
+
31
+ from . import debug
32
+ from .adapters import ModelAdapter
33
+ from .reducers import ForwardContext, Reducer
34
+ from .storage import RunWriter
35
+
36
+
37
+ @dataclass(frozen=True)
38
+ class _GenerationSample:
39
+ """The one document a generation belongs to, in the shape `debug.note_samples` reads.
40
+
41
+ A generate task runs one forward per decoded token and there is no batch, so this carries the same three fields a `ForwardContext` does and a fixed `batch_row` of 0.
42
+ Building a whole `ForwardContext` here is not possible: its positions are only known once the step has happened.
43
+ """
44
+
45
+ task_name: str
46
+ doc_id: int
47
+ choice_idx: int
48
+ batch_row: int = 0
49
+
50
+
51
+ class Recorder:
52
+ """Owns the hooks, the current-document context and the reducer calls.
53
+
54
+ Kept separate from `backend.py` on purpose: the backend holds a copy of lm-eval's `_loglikelihood_tokens`, and mixing our own logic into that file would make the source-hash check and the upstream diff harder to read.
55
+
56
+ Args:
57
+ hooks: 선택적 HookSpec 목록. 입력·출력에서 토큰별 분석 지표를 추출한다.
58
+ pass_name: evaluation 또는 collection. 이 pass에 해당하는 hook만 부착한다.
59
+ adapter: resolved module locations for this model.
60
+ reducers: the signal set for this run, in call order.
61
+ writer: where rows and tensors go.
62
+ tokenizer: used to render the `steps` table's token strings.
63
+ write_steps: emit `steps` rows.
64
+ False on the collection pass, where the first pass already recorded them for every document; writing them again would duplicate the rows of the collected documents.
65
+ skip_recorded: skip documents the run directory already holds.
66
+ A run that died leaves complete shards behind, and re-recording those documents on restart would duplicate their rows.
67
+ Scoring still re-runs them - lm-eval needs every request to produce a score - so this skips the recording, not the forward pass.
68
+
69
+ Example:
70
+ >>> rec = Recorder(adapter, [lens, sim], writer, tokenizer) # doctest: +SKIP
71
+ >>> with rec.session():
72
+ ... rec.expect_loglikelihood([ctx]) # says what the next forward covers
73
+ ... logits = model(input_ids).logits # hooks fire, reducers run
74
+ ... rec.flush() # rows are written
75
+ """
76
+
77
+ def __init__(
78
+ self,
79
+ adapter: ModelAdapter,
80
+ reducers: Sequence[Reducer],
81
+ writer: RunWriter,
82
+ tokenizer: Any,
83
+ write_steps: bool = True,
84
+ skip_recorded: bool = True,
85
+ hooks: Sequence[Any] = (),
86
+ pass_name: str = "evaluation",
87
+ ) -> None:
88
+ self.adapter = adapter
89
+ self.reducers = list(reducers)
90
+ self.writer = writer
91
+ self.tokenizer = tokenizer
92
+ self.write_steps = write_steps
93
+ self.already_recorded = set(writer.already_recorded) if skip_recorded else set()
94
+ # A collection pass re-runs recorded documents on purpose; a repeated collection
95
+ # replaces tensor files but must not append table rows (attn_norm) a second time.
96
+ self._rows_on_disk: dict[str, set[tuple[str, int, int]]] = {}
97
+ if not skip_recorded:
98
+ from .storage import existing_doc_keys
99
+ self._rows_on_disk = {r.table: existing_doc_keys(writer.run_dir, r.table)
100
+ for r in reducers if r.table}
101
+ self.skipped = 0
102
+
103
+ self._handles: list[Any] = []
104
+ # What the forward pass now in flight is for.
105
+ # None means "not tracing", which is how anything outside a declared request - a batch-size probe, or a document already recorded by an earlier run - produces no rows.
106
+ self._plan: dict[str, Any] | None = None
107
+ # Every context we built since the last flush, i.e. one document's worth for generate and one batch's worth for loglikelihood.
108
+ self._recorded: list[ForwardContext] = []
109
+ # Raw tensors caught during the current forward, cleared as soon as the reducers have consumed them.
110
+ self._attention: dict[int, torch.Tensor] = {}
111
+ self._values: dict[int, torch.Tensor] = {}
112
+ from .hooks import HookRuntime, load_hooks
113
+ self.custom = HookRuntime(self, load_hooks(None, {}, hooks), pass_name)
114
+
115
+ # -- what the reducers ask for -------------------------------------------------
116
+
117
+ @property
118
+ def wants_attention(self) -> bool:
119
+ """True when some reducer consumes attention weights."""
120
+ return any(r.source == "attention" for r in self.reducers)
121
+
122
+ @property
123
+ def wants_values(self) -> bool:
124
+ """True when some reducer consumes value vectors."""
125
+ return any(r.source == "value" for r in self.reducers)
126
+
127
+ # -- hook lifetime -------------------------------------------------------------
128
+
129
+ @contextmanager
130
+ def session(self) -> Iterator["Recorder"]:
131
+ """Register hooks for a block of evaluation, then always remove them.
132
+
133
+ A hook left behind keeps firing during unrelated work and keeps references to tensors that should have been freed, so the pairing is enforced by the context manager rather than by discipline.
134
+ """
135
+ try:
136
+ self.register_hooks()
137
+ yield self
138
+ except BaseException:
139
+ self.custom.drain()
140
+ self._recorded.clear()
141
+ raise
142
+ finally:
143
+ self.remove_hooks()
144
+
145
+ def register_hooks(self) -> None:
146
+ """등록 실패 시 부분적으로 부착된 handle도 회수한다."""
147
+ try:
148
+ self._register_hooks()
149
+ except BaseException:
150
+ self.remove_hooks()
151
+ raise
152
+
153
+ def _register_hooks(self) -> None:
154
+ if self._handles:
155
+ return
156
+ decoder = self.adapter.decoder
157
+ self._handles.append(
158
+ decoder.register_forward_pre_hook(self._request_hidden_states, with_kwargs=True)
159
+ )
160
+ self._handles.append(decoder.register_forward_hook(self._on_decoder_output))
161
+
162
+ if self.wants_attention:
163
+ for block_idx, module in enumerate(self.adapter.attn_modules):
164
+ if module is not None:
165
+ self._handles.append(
166
+ module.register_forward_hook(self._make_attention_hook(block_idx))
167
+ )
168
+ if self.wants_values:
169
+ for block_idx, module in enumerate(self.adapter.v_projs):
170
+ if module is not None:
171
+ self._handles.append(
172
+ module.register_forward_hook(self._make_value_hook(block_idx))
173
+ )
174
+
175
+ self.custom.register()
176
+ root = self.adapter.root_model
177
+ if root is not None:
178
+ self._handles.append(root.register_forward_pre_hook(self.custom.note_root, with_kwargs=True))
179
+ self._handles.append(root.register_forward_hook(lambda *args: self.custom.end()))
180
+
181
+ def remove_hooks(self) -> None:
182
+ self.custom.close()
183
+ self._attention.clear()
184
+ self._values.clear()
185
+ self._plan = None
186
+ for handle in self._handles:
187
+ handle.remove()
188
+ self._handles.clear()
189
+
190
+ # -- hooks ---------------------------------------------------------------------
191
+
192
+ def _request_hidden_states(self, module, args, kwargs): # noqa: ANN001 - torch signature
193
+ """Ask the decoder for its per-layer hidden states, one forward at a time."""
194
+ self.custom.begin(args, kwargs)
195
+ if self._plan is None:
196
+ return None
197
+ if any(r.source == "residual" for r in self.reducers):
198
+ kwargs["output_hidden_states"] = True
199
+ return args, kwargs
200
+
201
+ def _on_decoder_output(self, module, args, output): # noqa: ANN001 - torch signature
202
+ """The forward pass is complete: build the context and run the reducers.
203
+
204
+ Inner modules finish before their parent, so the attention and value hooks have already filled their buffers by the time this runs.
205
+ """
206
+ if self._plan is None:
207
+ return
208
+ hidden_states = getattr(output, "hidden_states", None)
209
+ needs_residual = any(r.source == "residual" for r in self.reducers)
210
+ if needs_residual:
211
+ if hidden_states is None:
212
+ raise RuntimeError(
213
+ "the decoder returned no hidden_states. The forward pre-hook that sets "
214
+ "output_hidden_states=True did not take effect; check the adapter's "
215
+ "decoder path for this architecture."
216
+ )
217
+ if len(hidden_states) != self.adapter.n_residual:
218
+ raise ValueError(
219
+ f"expected {self.adapter.n_residual} hidden states (L+1) but got "
220
+ f"{len(hidden_states)}; the adapter's block list does not match the model"
221
+ )
222
+ sequence_length = hidden_states[0].shape[1]
223
+ else:
224
+ # Attention/value/custom hooks need token positions, not all residuals.
225
+ last_hidden = getattr(output, "last_hidden_state", None)
226
+ if last_hidden is None:
227
+ raise RuntimeError("signal collection requires decoder.last_hidden_state for token positions")
228
+ sequence_length = last_hidden.shape[1]
229
+ hidden_states = ()
230
+ # The logit lens decodes through the model's *own* final norm and head, so with a
231
+ # module tracer attached those calls are module entries nested inside this hook.
232
+ # The phase is what keeps them from being read as the model's own work.
233
+ # Snapshot before reducers advance generation and invoke auxiliary head/norm calls.
234
+ self.custom.contexts(sequence_length)
235
+ self.custom.suspended = True
236
+ try:
237
+ with debug.phase("signal_collection"):
238
+ for ctx in self._contexts_for_forward(sequence_length):
239
+ self._run_reducers(ctx, hidden_states)
240
+ self._recorded.append(ctx)
241
+ finally:
242
+ self.custom.suspended = False
243
+ if self.adapter.root_model is None:
244
+ self.custom.end()
245
+ self._attention.clear()
246
+ self._values.clear()
247
+
248
+ def _make_attention_hook(self, block_idx: int):
249
+ """Grab the (batch, heads, q, k) attention probabilities of one block."""
250
+
251
+ def hook(module, args, output): # noqa: ANN001 - torch signature
252
+ if self._plan is None:
253
+ return
254
+ weights = _find_attention_weights(output)
255
+ if weights is None:
256
+ raise RuntimeError(
257
+ f"block {block_idx}: the attention module returned no weight matrix. "
258
+ "Attention capture needs attn_implementation='eager'; FlashAttention "
259
+ "and SDPA never materialise the probabilities."
260
+ )
261
+ self._attention[block_idx] = weights.detach()
262
+
263
+ return hook
264
+
265
+ def _make_value_hook(self, block_idx: int):
266
+ """Grab the value projection output of one block."""
267
+
268
+ def hook(module, args, output): # noqa: ANN001 - torch signature
269
+ if self._plan is None:
270
+ return
271
+ tensor = output[0] if isinstance(output, tuple) else output
272
+ # How to read values out of this projection is architecture specific, so the adapter owns it.
273
+ self._values[block_idx] = self.adapter.extract_value(tensor.detach(), block_idx)
274
+
275
+ return hook
276
+
277
+ # -- declaring what a forward pass is for --------------------------------------
278
+
279
+ def expect_loglikelihood(self, contexts: Sequence[ForwardContext]) -> None:
280
+ """Declare the documents of the next loglikelihood forward, one per batch row.
281
+
282
+ This is the step that keeps signals attached to the right document. lm-eval drops the `Instance` (and with it `doc_id`) before the model sees a request, and then reorders requests by length; `backend.py` re-establishes the link and states it here, right before the call.
283
+
284
+ Documents already on disk from an earlier, interrupted run are dropped here.
285
+ The forward pass still happens - lm-eval is scoring them - but nothing is recorded, so the shards gain no duplicates.
286
+ """
287
+ # The tracer is told about every context, not just the ones that will be recorded:
288
+ # a document already on disk is still scored, and a forward that is not recorded is
289
+ # exactly as able to run out of memory as one that is.
290
+ debug.note_samples(contexts)
291
+ wanted = [ctx for ctx in contexts if ctx.key not in self.already_recorded]
292
+ self.skipped += len(contexts) - len(wanted)
293
+ self._plan = {"kind": "loglikelihood", "contexts": wanted} if wanted else None
294
+
295
+ def expect_generation(
296
+ self, task_name: str, doc_id: int, prompt_length: int, choice_idx: int = 0
297
+ ) -> None:
298
+ """Declare that the next `model.generate()` belongs to one document.
299
+
300
+ Generation runs one forward per token, so the contexts are built as the steps happen rather than up front: `step 0` is the last prefill position, and every decode step after that adds one.
301
+
302
+ Args:
303
+ prompt_length: number of prompt tokens, needed to turn a decoding step into an absolute sequence position.
304
+ """
305
+ debug.note_samples(
306
+ [_GenerationSample(task_name, doc_id, choice_idx)]
307
+ )
308
+ if (task_name, doc_id, choice_idx) in self.already_recorded:
309
+ self.skipped += 1
310
+ self._plan = None
311
+ return
312
+ self._plan = {
313
+ "kind": "generate",
314
+ "task_name": task_name,
315
+ "doc_id": doc_id,
316
+ "choice_idx": choice_idx,
317
+ "prompt_length": prompt_length,
318
+ "next_step": 0,
319
+ }
320
+
321
+ def set_prompt_length(self, prompt_length: int) -> None:
322
+ """Tell an in-flight generation how long its prompt is.
323
+
324
+ Only known once lm-eval has tokenised and padded the prompt, which happens after `expect_generation` and before the first forward pass.
325
+ Without it a decoding step cannot be turned into an absolute position.
326
+ """
327
+ if self._plan is not None and self._plan.get("kind") == "generate":
328
+ self._plan["prompt_length"] = int(prompt_length)
329
+
330
+ def _contexts_for_forward(self, seq_len: int) -> list[ForwardContext]:
331
+ """Build the contexts describing the forward pass that just finished."""
332
+ plan = self._plan
333
+ assert plan is not None
334
+ if plan["kind"] == "loglikelihood":
335
+ return list(plan["contexts"])
336
+
337
+ # generate: prefill carries the whole prompt and we score its last position; every later forward is a single decoded token.
338
+ step = plan["next_step"]
339
+ plan["next_step"] = step + 1
340
+ seq_index = seq_len - 1
341
+ position = plan["prompt_length"] - 1 + step
342
+ return [
343
+ ForwardContext(
344
+ task_name=plan["task_name"],
345
+ doc_id=plan["doc_id"],
346
+ choice_idx=plan["choice_idx"],
347
+ steps=[step],
348
+ positions=[position],
349
+ target_token_ids=None, # a generate task has no gold token to rank
350
+ n_residual=self.adapter.n_residual,
351
+ n_blocks=self.adapter.n_blocks,
352
+ task_kind="generate",
353
+ batch_row=0,
354
+ seq_indices=[seq_index],
355
+ input_offset=(max(0, plan["original_prompt_length"] - plan["prompt_length"])
356
+ if "original_prompt_length" in plan else None),
357
+ )
358
+ ]
359
+
360
+ def _run_reducers(self, ctx: ForwardContext, hidden_states: Sequence[torch.Tensor]) -> None:
361
+ """Feed one document's slice of this forward pass to every reducer."""
362
+ for reducer in self.reducers:
363
+ if reducer.source == "residual":
364
+ for layer, tensor in enumerate(hidden_states):
365
+ reducer.update(layer, _slice_positions(tensor, ctx), ctx)
366
+ elif reducer.source == "attention":
367
+ for block, tensor in sorted(self._attention.items()):
368
+ reducer.update(block, _slice_attention(tensor, ctx), ctx)
369
+ elif reducer.source == "value":
370
+ for block, tensor in sorted(self._values.items()):
371
+ reducer.update(block, _slice_positions(tensor, ctx), ctx)
372
+ reducer.end_forward(ctx)
373
+ ctx.shared.clear()
374
+
375
+ # -- finishing documents -------------------------------------------------------
376
+
377
+ def set_generated_tokens(self, token_ids: Sequence[int]) -> None:
378
+ """Attach the tokens a generate call actually emitted to their steps.
379
+
380
+ Only known once generation has finished, and needed because sampling can make the emitted token differ from the last layer's top-1.
381
+
382
+ Example:
383
+ >>> rec.set_generated_tokens([1820, 4320, 128009]) # doctest: +SKIP
384
+ """
385
+ for ctx in self._recorded:
386
+ if ctx.task_kind != "generate":
387
+ continue
388
+ step = ctx.steps[0]
389
+ ctx.step_token_ids = [int(token_ids[step])] if step < len(token_ids) else [None]
390
+
391
+ def flush(self) -> int:
392
+ """Finalize the reducers and write everything recorded since the last flush.
393
+
394
+ Returns:
395
+ How many documents were written.
396
+ """
397
+ contexts = self._recorded
398
+ self._recorded = []
399
+ self._plan = None
400
+ if not contexts:
401
+ return 0
402
+
403
+ documents = _group_by_document(contexts)
404
+ rows_by_table: dict[str, list[dict[str, Any]]] = self.custom.drain()
405
+ if self.write_steps:
406
+ rows_by_table["steps"] = []
407
+ for group in documents:
408
+ rows_by_table["steps"].extend(self.steps_rows(group))
409
+
410
+ for reducer in self.reducers:
411
+ rows = reducer.finalize()
412
+ if rows and reducer.table:
413
+ on_disk = self._rows_on_disk.get(reducer.table)
414
+ if on_disk:
415
+ rows = [row for row in rows
416
+ if (row["task_name"], int(row["doc_id"]), int(row["choice_idx"])) not in on_disk]
417
+ rows_by_table.setdefault(reducer.table, []).extend(rows)
418
+ self._write_tensors(reducer, documents)
419
+
420
+ self.writer.write_document(rows_by_table, documents=len(documents))
421
+ return len(documents)
422
+
423
+ def steps_rows(self, contexts: Sequence[ForwardContext]) -> list[dict[str, Any]]:
424
+ """Build the `steps` rows for one document, across all of its forwards.
425
+
426
+ The `steps` table cannot be replaced by the layer-L rows of `signals`: for loglikelihood the scored token is the gold continuation token rather than the model's prediction, and with sampling on, a generated token can differ from the last layer's top-1.
427
+
428
+ Example:
429
+ >>> rec.steps_rows([ctx]) # doctest: +SKIP
430
+ [{'task_name': 'xnli_ko', 'doc_id': 7, 'choice_idx': 1, 'step': 0,
431
+ 'token_id': 9891, 'token': 'Ġyes', 'position': 41}]
432
+ """
433
+ rows = []
434
+ for ctx in contexts:
435
+ token_ids = ctx.step_token_ids or []
436
+ for index, step in enumerate(ctx.steps):
437
+ token_id = token_ids[index] if index < len(token_ids) else None
438
+ token_id = None if token_id is None else int(token_id)
439
+ rows.append(
440
+ {
441
+ "task_name": ctx.task_name,
442
+ "doc_id": ctx.doc_id,
443
+ "choice_idx": ctx.choice_idx,
444
+ "step": step,
445
+ "token_id": token_id,
446
+ "token": self._decode_token(token_id),
447
+ "position": ctx.positions[index],
448
+ }
449
+ )
450
+ return rows
451
+
452
+ def _decode_token(self, token_id: int | None) -> str | None:
453
+ if token_id is None:
454
+ return None
455
+ try:
456
+ return str(self.tokenizer.convert_ids_to_tokens(token_id))
457
+ except Exception:
458
+ return str(self.tokenizer.decode([token_id]))
459
+
460
+ def _write_tensors(
461
+ self, reducer: Reducer, documents: Sequence[Sequence[ForwardContext]]
462
+ ) -> None:
463
+ """Persist a reducer's opt-in tensor output, if it produced any.
464
+
465
+ Tensor dumps are one file per document, so they are only defined for an unbatched forward - which is exactly how the collection path runs.
466
+ """
467
+ tensors = reducer.pop_tensors()
468
+ if not tensors:
469
+ return
470
+ if len(documents) != 1:
471
+ raise RuntimeError(
472
+ f"{reducer.name} writes one tensor file per document, so it requires "
473
+ f"batch size 1, but this flush covered {len(documents)} documents"
474
+ )
475
+ ctx = documents[0][0]
476
+ store = self.writer.attention if reducer.source == "attention" else self.writer.raw_hidden
477
+ meta = {
478
+ "reducer": reducer.name,
479
+ "version": str(reducer.version),
480
+ "task_kind": ctx.task_kind,
481
+ }
482
+ store.save(ctx.task_name, ctx.doc_id, ctx.choice_idx, tensors, meta=meta)
483
+
484
+
485
+ # --------------------------------------------------------------------------
486
+ # Helpers
487
+ # --------------------------------------------------------------------------
488
+
489
+
490
+ def _group_by_document(
491
+ contexts: Sequence[ForwardContext],
492
+ ) -> list[list[ForwardContext]]:
493
+ """Group forward contexts by (task_name, doc_id, choice_idx), keeping order.
494
+
495
+ A loglikelihood batch gives one context per document; a generation gives many contexts - one per decoding step - that all belong to one document.
496
+
497
+ Example:
498
+ >>> a = ForwardContext("t", 1, 0, [0], [3], None, 2, 1)
499
+ >>> b = ForwardContext("t", 1, 0, [1], [4], None, 2, 1)
500
+ >>> c = ForwardContext("t", 2, 0, [0], [3], None, 2, 1)
501
+ >>> [len(g) for g in _group_by_document([a, b, c])]
502
+ [2, 1]
503
+ """
504
+ grouped: dict[tuple[str, int, int], list[ForwardContext]] = {}
505
+ for ctx in contexts:
506
+ grouped.setdefault(ctx.key, []).append(ctx)
507
+ return list(grouped.values())
508
+
509
+
510
+ def _slice_positions(tensor: torch.Tensor, ctx: ForwardContext) -> torch.Tensor:
511
+ """Cut a (batch, seq, feature) tensor down to one document's scored positions.
512
+
513
+ Right padding lives at the end of the sequence axis, so reading the tensor's last index instead of the document's own would silently pick up a pad token.
514
+ Selecting explicit indices avoids that entirely.
515
+
516
+ Example:
517
+ >>> x = torch.arange(2 * 4 * 3).reshape(2, 4, 3)
518
+ >>> ctx = ForwardContext("t", 0, 0, [0], [2], None, 2, 1, batch_row=1)
519
+ >>> _slice_positions(x, ctx).tolist()
520
+ [[18, 19, 20]]
521
+ """
522
+ index = torch.as_tensor(ctx.seq_indices, device=tensor.device, dtype=torch.long)
523
+ return tensor[ctx.batch_row].index_select(0, index)
524
+
525
+
526
+ def _slice_attention(tensor: torch.Tensor, ctx: ForwardContext) -> torch.Tensor:
527
+ """Cut a (batch, heads, q, k) attention tensor down to (n_pos, heads, k).
528
+
529
+ Only the query rows we score are kept; the full q x k map is what makes attention unaffordable to store.
530
+
531
+ Example:
532
+ >>> attn = torch.zeros(1, 4, 6, 6)
533
+ >>> ctx = ForwardContext("t", 0, 0, [0], [5], None, 2, 1)
534
+ >>> _slice_attention(attn, ctx).shape
535
+ torch.Size([1, 4, 6])
536
+ """
537
+ index = torch.as_tensor(ctx.seq_indices, device=tensor.device, dtype=torch.long)
538
+ rows = tensor[ctx.batch_row].index_select(1, index) # (heads, n_pos, k)
539
+ return rows.transpose(0, 1) # (n_pos, heads, k)
540
+
541
+
542
+ def _find_attention_weights(output: Any) -> torch.Tensor | None:
543
+ """Pick the attention probability matrix out of whatever the module returned.
544
+
545
+ HF attention modules return `(attn_output, attn_weights)` and sometimes a third element, with `attn_weights` set to None unless the implementation materialises it.
546
+ The probabilities are the only 4-dimensional element, so we look for that rather than for a tuple position that shifts between transformers versions.
547
+
548
+ Example:
549
+ >>> _find_attention_weights((torch.zeros(1, 3, 8), torch.zeros(1, 2, 3, 3))).shape
550
+ torch.Size([1, 2, 3, 3])
551
+ """
552
+ candidates = output if isinstance(output, (tuple, list)) else [output]
553
+ for item in candidates:
554
+ if isinstance(item, torch.Tensor) and item.dim() == 4:
555
+ return item
556
+ return None