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/__init__.py +27 -0
- evalmetry/adapters.py +576 -0
- evalmetry/backend.py +942 -0
- evalmetry/benchmarks.py +261 -0
- evalmetry/debug.py +1421 -0
- evalmetry/hooks.py +782 -0
- evalmetry/judges.py +273 -0
- evalmetry/main.py +1504 -0
- evalmetry/models.py +209 -0
- evalmetry/module_stats.py +1084 -0
- evalmetry/recorder.py +556 -0
- evalmetry/reducers.py +578 -0
- evalmetry/report.py +2193 -0
- evalmetry/storage.py +1546 -0
- evalmetry-1.0.0.dist-info/METADATA +94 -0
- evalmetry-1.0.0.dist-info/RECORD +20 -0
- evalmetry-1.0.0.dist-info/WHEEL +5 -0
- evalmetry-1.0.0.dist-info/entry_points.txt +2 -0
- evalmetry-1.0.0.dist-info/licenses/LICENSE +21 -0
- evalmetry-1.0.0.dist-info/top_level.txt +1 -0
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
|