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
|
@@ -0,0 +1,1084 @@
|
|
|
1
|
+
"""Sample-level module statistics and reproducible dataset aggregation.
|
|
2
|
+
|
|
3
|
+
Execution traces answer where a call stopped. This file instead stores one row per
|
|
4
|
+
physical batch row, module invocation and tensor, then pools those observations by
|
|
5
|
+
document before comparing documents. SQLite keeps collection and aggregation bounded
|
|
6
|
+
in host memory; Parquet exports are for downstream analysis. No activations are kept.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import itertools
|
|
12
|
+
import json
|
|
13
|
+
import math
|
|
14
|
+
import os
|
|
15
|
+
import re
|
|
16
|
+
import sqlite3
|
|
17
|
+
import uuid
|
|
18
|
+
from datetime import datetime, timezone
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
from typing import Any, Iterable
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
SCHEMA_VERSION = 5
|
|
24
|
+
CHUNK_ELEMENTS = 65536
|
|
25
|
+
GROUP = ("task_name", "module", "io", "tensor_path", "dtype", "stage", "layout")
|
|
26
|
+
# One axis index of one tensor is its own population, so it joins the group key
|
|
27
|
+
# rather than being pooled into it. `axis_role` is determined by the rest of the key
|
|
28
|
+
# and travels with it so the meaning of an index is readable without a join.
|
|
29
|
+
AXIS_GROUP = GROUP + ("axis", "axis_role", "axis_index")
|
|
30
|
+
|
|
31
|
+
# Boundaries whose tensors are [batch * width, features] in the model's own token
|
|
32
|
+
# order, keyed by the *parent* module's class and the attribute holding the child.
|
|
33
|
+
# The child's own class is not enough to identify one: a transformers 4.57 router is
|
|
34
|
+
# a bare `nn.Linear`, and 5.16 fuses every expert into one module.
|
|
35
|
+
#
|
|
36
|
+
# Read from the installed implementations, not inferred. `Qwen3MoeSparseMoeBlock`
|
|
37
|
+
# and `MixtralSparseMoeBlock` both flatten with `hidden_states.view(-1, hidden_dim)`
|
|
38
|
+
# before calling `gate` and `experts`, and reshape the experts' result back to
|
|
39
|
+
# [batch, sequence, hidden]. Row `r` of either boundary is therefore document row
|
|
40
|
+
# `r // width` at position `r % width`, and the expert-major gather happens strictly
|
|
41
|
+
# inside the experts module. In 4.57 the experts are an `nn.ModuleList`, so each
|
|
42
|
+
# expert's parent is that list and none of them is registered here: a single expert
|
|
43
|
+
# sees only the tokens routed to it, in routing order.
|
|
44
|
+
FLATTENED_BOUNDARIES = {
|
|
45
|
+
"Qwen3MoeSparseMoeBlock.gate", "Qwen3MoeSparseMoeBlock.experts",
|
|
46
|
+
"MixtralSparseMoeBlock.gate", "MixtralSparseMoeBlock.experts",
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
# Registered routers whose own output carries the dispatch its parent block used, and
|
|
50
|
+
# what each of their outputs means. Read from the installed implementations, not
|
|
51
|
+
# inferred: `Qwen3MoeTopKRouter.forward` and `MixtralTopKRouter.forward` both return
|
|
52
|
+
# `(router_logits, router_scores, router_indices)`, and both blocks hand exactly those
|
|
53
|
+
# scores and indices to `experts`. So the selection here is observed, not reconstructed.
|
|
54
|
+
#
|
|
55
|
+
# The logits also appear alone: in transformers 4.57 the same `gate` attribute is a bare
|
|
56
|
+
# `nn.Linear` returning one tensor, and the softmax and top-k are taken in the parent
|
|
57
|
+
# block, where no registered boundary can see them. Such a run records the logits per
|
|
58
|
+
# expert and says routing was unavailable, rather than recomputing a selection whose
|
|
59
|
+
# dtype, normalisation and k are the implementation's and not observable at the boundary.
|
|
60
|
+
_TOP_K_ROUTER = {
|
|
61
|
+
"logits": ("output.0", "output"),
|
|
62
|
+
"weights": "output.1",
|
|
63
|
+
"experts": "output.2",
|
|
64
|
+
"selection": "observed",
|
|
65
|
+
}
|
|
66
|
+
ROUTER_BOUNDARIES = {
|
|
67
|
+
"Qwen3MoeSparseMoeBlock.gate": _TOP_K_ROUTER,
|
|
68
|
+
"MixtralSparseMoeBlock.gate": _TOP_K_ROUTER,
|
|
69
|
+
}
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _transformers_version() -> str | None:
|
|
73
|
+
"""Identify the implementation the registry was matched against, if it is there."""
|
|
74
|
+
try:
|
|
75
|
+
import transformers
|
|
76
|
+
except ImportError:
|
|
77
|
+
return None
|
|
78
|
+
return getattr(transformers, "__version__", None)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def merge_moments(a: dict, b: dict) -> dict:
|
|
82
|
+
"""Combine finite-value moments using the parallel variance formula.
|
|
83
|
+
|
|
84
|
+
Counts include non-finite values, while mean/M2/min/max describe finite values
|
|
85
|
+
only. Population variance is M2 / finite_count, never an average of variances.
|
|
86
|
+
"""
|
|
87
|
+
out = {k: a.get(k, 0) + b.get(k, 0)
|
|
88
|
+
for k in ("count", "finite_count", "nan", "posinf", "neginf")}
|
|
89
|
+
n, m = a.get("finite_count", 0), b.get("finite_count", 0)
|
|
90
|
+
if not n or not m:
|
|
91
|
+
source = a if n else b
|
|
92
|
+
out.update({k: source.get(k) for k in ("mean", "m2", "min", "max")})
|
|
93
|
+
else:
|
|
94
|
+
delta = b["mean"] - a["mean"]
|
|
95
|
+
out.update(mean=a["mean"] + delta * (m / (n + m)),
|
|
96
|
+
m2=a["m2"] + b["m2"] + delta * delta * (n * m / (n + m)),
|
|
97
|
+
min=min(a["min"], b["min"]), max=max(a["max"], b["max"]))
|
|
98
|
+
return out
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def tensor_moments(tensor: Any, chunk_elements: int | None = None) -> dict:
|
|
102
|
+
"""Reduce a floating tensor in bounded float64 chunks without keeping its graph.
|
|
103
|
+
|
|
104
|
+
Float64 prevents the tracer from overflowing when squaring large float32/bf16
|
|
105
|
+
activations. The temporary reduction buffers are bounded; flattening a strided
|
|
106
|
+
input can still require a contiguous copy of that sample, so the chunk size alone
|
|
107
|
+
does not bound the memory a reduction takes.
|
|
108
|
+
"""
|
|
109
|
+
import torch
|
|
110
|
+
|
|
111
|
+
total: dict = {}
|
|
112
|
+
for chunk in tensor.detach().reshape(-1).split(chunk_elements or CHUNK_ELEMENTS):
|
|
113
|
+
values = chunk.to(dtype=torch.float64)
|
|
114
|
+
counts = torch.stack([torch.isnan(values).sum(), torch.isposinf(values).sum(),
|
|
115
|
+
torch.isneginf(values).sum()]).cpu().tolist()
|
|
116
|
+
finite = values[torch.isfinite(values)]
|
|
117
|
+
part = dict(count=chunk.numel(), finite_count=finite.numel(),
|
|
118
|
+
nan=counts[0], posinf=counts[1], neginf=counts[2])
|
|
119
|
+
if finite.numel():
|
|
120
|
+
mean = finite.mean()
|
|
121
|
+
packed = torch.stack([mean, ((finite - mean) ** 2).sum(),
|
|
122
|
+
finite.min(), finite.max()]).cpu().tolist()
|
|
123
|
+
part.update(zip(("mean", "m2", "min", "max"), packed))
|
|
124
|
+
total = merge_moments(total, part)
|
|
125
|
+
return merge_moments({}, total)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def _output_tensor(value: Any, spec: str) -> Any:
|
|
129
|
+
"""The tensor a registry path names, or None when this output has no such member.
|
|
130
|
+
|
|
131
|
+
`output` is the returned value itself, `output.N` its N-th member. Anything else at
|
|
132
|
+
that position - a different container, a version that returns fewer values - is None
|
|
133
|
+
rather than a guess at which member was meant.
|
|
134
|
+
"""
|
|
135
|
+
import torch
|
|
136
|
+
|
|
137
|
+
if spec == "output":
|
|
138
|
+
return value if isinstance(value, torch.Tensor) else None
|
|
139
|
+
index = int(spec.rpartition(".")[2])
|
|
140
|
+
if isinstance(value, (tuple, list)) and index < len(value):
|
|
141
|
+
member = value[index]
|
|
142
|
+
return member if isinstance(member, torch.Tensor) else None
|
|
143
|
+
return None
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _tensors(value: Any, path: str, depth: int = 0):
|
|
147
|
+
"""Walk tensor containers without truncating layers or retaining cache objects.
|
|
148
|
+
|
|
149
|
+
KV caches are deliberately opaque: re-counting cached tokens on every decode
|
|
150
|
+
step is a different population from newly computed module activations.
|
|
151
|
+
"""
|
|
152
|
+
import torch
|
|
153
|
+
|
|
154
|
+
if isinstance(value, torch.Tensor):
|
|
155
|
+
yield path, value
|
|
156
|
+
elif depth < 8 and isinstance(value, dict):
|
|
157
|
+
for key, child in value.items():
|
|
158
|
+
yield from _tensors(child, f"{path}.{key}", depth + 1)
|
|
159
|
+
elif depth < 8 and isinstance(value, (tuple, list)):
|
|
160
|
+
for index, child in enumerate(value):
|
|
161
|
+
yield from _tensors(child, f"{path}.{index}", depth + 1)
|
|
162
|
+
elif value is not None and not isinstance(value, (bool, int, float, str)):
|
|
163
|
+
yield path, None
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
class ModuleStatistics:
|
|
167
|
+
"""Persist sample observations in a new session, never mixing resumed attempts.
|
|
168
|
+
|
|
169
|
+
Supported activation layout is [batch, current_sequence, features]. Attention
|
|
170
|
+
matrices use [batch, heads, query, key]. Other layouts are counted as excluded,
|
|
171
|
+
so a shared tensor is not silently assigned to each document. The backend must
|
|
172
|
+
supply true input lengths; scoring positions alone are not a padding mask.
|
|
173
|
+
"""
|
|
174
|
+
|
|
175
|
+
def __init__(self, trace_path: str, config: dict, chunk_elements: int | None = None):
|
|
176
|
+
base = Path(trace_path)
|
|
177
|
+
# A reduction setting, not part of `config`: two sessions that differ only here
|
|
178
|
+
# must hold the same numbers, and keeping it out of `config` lets them compare so.
|
|
179
|
+
self.chunk_elements = int(chunk_elements or CHUNK_ELEMENTS)
|
|
180
|
+
if self.chunk_elements < 1:
|
|
181
|
+
raise ValueError("chunk_elements must be at least 1")
|
|
182
|
+
session = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S%f") + "-" + uuid.uuid4().hex[:8]
|
|
183
|
+
self.directory = base.parent / "module_stats" / base.stem / session
|
|
184
|
+
self.directory.mkdir(parents=True)
|
|
185
|
+
self.path = self.directory / "statistics.sqlite"
|
|
186
|
+
self._configure_observations(config)
|
|
187
|
+
self._create_database()
|
|
188
|
+
self._write_session_metadata(config, base.stem, session)
|
|
189
|
+
|
|
190
|
+
# 아래 상태는 forward마다 갱신되거나 모듈 등록 때 채워진다.
|
|
191
|
+
# 저장 스키마·설정과 분리해 현재 관측 상태를 한곳에서 확인한다.
|
|
192
|
+
self.contexts: list[Any] = []
|
|
193
|
+
self.next_step = 0
|
|
194
|
+
self.forward = 0
|
|
195
|
+
self.stage = "unknown"
|
|
196
|
+
self.step = 0
|
|
197
|
+
self.lengths: list[int] | None = None
|
|
198
|
+
self.width = 0
|
|
199
|
+
self.batch = 0
|
|
200
|
+
self.flattened: dict[str, str | None] = {}
|
|
201
|
+
self.extremes: dict[str, int] = {}
|
|
202
|
+
self.axes: dict[str, str | None] = {}
|
|
203
|
+
self.routers: dict[str, dict | None] = {}
|
|
204
|
+
self.routing: dict[str, dict | None] = {}
|
|
205
|
+
self.closed = False
|
|
206
|
+
|
|
207
|
+
def _configure_observations(self, config: dict) -> None:
|
|
208
|
+
"""추가 관측의 옵션과 모듈 selector를 해석한다. 실제 경로는 등록 때 확정한다."""
|
|
209
|
+
# Extreme positions are opt-in twice over: a k, and the modules that record it.
|
|
210
|
+
self.extremes_k = int(config.get("extremes") or 0)
|
|
211
|
+
selector = config.get("extremes_modules")
|
|
212
|
+
self.extremes_pattern = re.compile(selector) if self.extremes_k and selector else None
|
|
213
|
+
# Per-axis statistics the same way: which axis to keep, and where to keep it.
|
|
214
|
+
# One row per index is a different order of output from one row per tensor.
|
|
215
|
+
self.axes_spec = config.get("axes") or None
|
|
216
|
+
axis_selector = config.get("axes_modules")
|
|
217
|
+
self.axes_pattern = (re.compile(axis_selector)
|
|
218
|
+
if self.axes_spec and axis_selector else None)
|
|
219
|
+
# Routing needs a selector alone: the k, the expert ids and the weights are the
|
|
220
|
+
# router's own output, so there is nothing left for a second flag to decide.
|
|
221
|
+
routing_selector = config.get("routing")
|
|
222
|
+
self.routing_pattern = re.compile(routing_selector) if routing_selector else None
|
|
223
|
+
|
|
224
|
+
def _create_database(self) -> None:
|
|
225
|
+
"""새 세션의 DB와 원시 관측 테이블을 만든다. 집계 테이블은 종료 때 생성한다."""
|
|
226
|
+
self.db = sqlite3.connect(self.path)
|
|
227
|
+
self.db.row_factory = sqlite3.Row
|
|
228
|
+
self.db.execute("PRAGMA temp_store=FILE")
|
|
229
|
+
self.db.executescript("""
|
|
230
|
+
CREATE TABLE metadata (key TEXT PRIMARY KEY, value TEXT NOT NULL);
|
|
231
|
+
CREATE TABLE modules (
|
|
232
|
+
module TEXT PRIMARY KEY, module_type TEXT, parent_type TEXT,
|
|
233
|
+
flattened_rule TEXT, extremes_k INTEGER, axis_stats TEXT,
|
|
234
|
+
router_rule TEXT, routing TEXT, router_experts INTEGER,
|
|
235
|
+
router_top_k INTEGER, registration_status TEXT,
|
|
236
|
+
registration_error TEXT, call_count INTEGER NOT NULL DEFAULT 0);
|
|
237
|
+
CREATE VIEW coverage AS SELECT m.*,
|
|
238
|
+
(SELECT count(*) FROM observations o WHERE o.module=m.module) AS observation_count,
|
|
239
|
+
coalesce((SELECT sum(e.count) FROM exclusions e WHERE e.module=m.module), 0) AS exclusion_count
|
|
240
|
+
FROM modules m;
|
|
241
|
+
CREATE TABLE forwards (forward INTEGER PRIMARY KEY, status TEXT NOT NULL);
|
|
242
|
+
CREATE TABLE observations (
|
|
243
|
+
forward INTEGER, call INTEGER, batch_row INTEGER,
|
|
244
|
+
task_name TEXT, doc_id INTEGER, choices TEXT,
|
|
245
|
+
module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
246
|
+
step INTEGER, layout TEXT, shape TEXT, source_shape TEXT,
|
|
247
|
+
count INTEGER, finite_count INTEGER,
|
|
248
|
+
nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
249
|
+
mean REAL, m2 REAL, min REAL, max REAL);
|
|
250
|
+
CREATE INDEX observations_module ON observations(module);
|
|
251
|
+
CREATE TABLE extremes (
|
|
252
|
+
forward INTEGER, call INTEGER, batch_row INTEGER,
|
|
253
|
+
task_name TEXT, doc_id INTEGER, choices TEXT,
|
|
254
|
+
module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
255
|
+
step INTEGER, layout TEXT, position INTEGER, feature INTEGER,
|
|
256
|
+
value REAL, abs_rank INTEGER);
|
|
257
|
+
CREATE INDEX extremes_module ON extremes(module);
|
|
258
|
+
CREATE TABLE axis_stats (
|
|
259
|
+
forward INTEGER, call INTEGER, batch_row INTEGER,
|
|
260
|
+
task_name TEXT, doc_id INTEGER, choices TEXT,
|
|
261
|
+
module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
262
|
+
step INTEGER, layout TEXT, axis TEXT, axis_role TEXT, axis_index INTEGER,
|
|
263
|
+
count INTEGER, finite_count INTEGER,
|
|
264
|
+
nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
265
|
+
mean REAL, m2 REAL, min REAL, max REAL);
|
|
266
|
+
CREATE INDEX axis_stats_module ON axis_stats(module);
|
|
267
|
+
CREATE TABLE routing (
|
|
268
|
+
forward INTEGER, call INTEGER, batch_row INTEGER,
|
|
269
|
+
task_name TEXT, doc_id INTEGER, choices TEXT,
|
|
270
|
+
module TEXT, stage TEXT, step INTEGER, layout TEXT, dtype TEXT,
|
|
271
|
+
position INTEGER, rank INTEGER, expert INTEGER,
|
|
272
|
+
weight REAL, weight_state TEXT, source TEXT, weight_scope TEXT);
|
|
273
|
+
CREATE INDEX routing_module ON routing(module);
|
|
274
|
+
CREATE TABLE exclusions (
|
|
275
|
+
module TEXT, io TEXT, tensor_path TEXT, reason TEXT, count INTEGER,
|
|
276
|
+
PRIMARY KEY (module, io, tensor_path, reason));
|
|
277
|
+
""")
|
|
278
|
+
|
|
279
|
+
def _write_session_metadata(self, config: dict, pass_name: str, session: str) -> None:
|
|
280
|
+
"""집계 기준과 좌표 의미를 DB에 기록해 모델 없이도 결과를 해석하게 한다.
|
|
281
|
+
|
|
282
|
+
아래 문자열은 출력 데이터의 일부다. 설명을 정리할 때도 저장되는 값과
|
|
283
|
+
키를 유지해야 기존 세션과 같은 기준으로 비교할 수 있다.
|
|
284
|
+
"""
|
|
285
|
+
self._meta("schema_version", SCHEMA_VERSION)
|
|
286
|
+
self._meta("config", config)
|
|
287
|
+
self._meta("pass", pass_name)
|
|
288
|
+
self._meta("session", session)
|
|
289
|
+
self._meta("status", "running")
|
|
290
|
+
self._meta("aggregation_status", "pending")
|
|
291
|
+
self._meta("attention_mask_policy", "exclude exact attention_mask path components, including nested containers; no positional or value inference")
|
|
292
|
+
self._meta("document_statistics", {
|
|
293
|
+
"weighting": "equal weight per document with finite values",
|
|
294
|
+
"sample_mean_std": "population std of document means: sqrt(M2/n)",
|
|
295
|
+
"sample_mean_m2": "sum of squared deviations of finite document means",
|
|
296
|
+
"sample_mean_sample_std": "sqrt(M2/(n-1)); null for n<2",
|
|
297
|
+
"sample_mean_sem": "sample std/sqrt(n); assumes independent sampled documents; null for n<2",
|
|
298
|
+
"n": "finite_documents; excludes documents containing only nonfinite values"})
|
|
299
|
+
self._meta("coverage_units", {
|
|
300
|
+
"call_count": "selected module pre-hook entries, all phases, including failed calls",
|
|
301
|
+
"observation_count": "stored document/tensor rows, input and output, including incomplete forwards",
|
|
302
|
+
"exclusion_count": "exclusion events; usually tensor visits, per-document for key-length or moment failures",
|
|
303
|
+
"durability": "targets and registration transitions committed immediately; calls at entry; observations after each observer"})
|
|
304
|
+
self._meta("flattened_layout", {
|
|
305
|
+
"rule": "parent module class + attribute, matched before the first hook",
|
|
306
|
+
"registered": sorted(FLATTENED_BOUNDARIES),
|
|
307
|
+
"row_index": "batch_row * width + position; right padding dropped per document",
|
|
308
|
+
"requires": "exactly two dimensions and first axis == batch * width",
|
|
309
|
+
"refuses": "unregistered boundaries, size mismatch from dropped or padded "
|
|
310
|
+
"tokens, and any expert-major child such as 4.57 experts.N",
|
|
311
|
+
"pooling": "router logit statistics reduce over the expert axis; they are "
|
|
312
|
+
"not per-expert observations",
|
|
313
|
+
"transformers": _transformers_version()})
|
|
314
|
+
self._meta("extremes", {
|
|
315
|
+
"k": self.extremes_k,
|
|
316
|
+
"modules": config.get("extremes_modules"),
|
|
317
|
+
"unit": "one row per rank, per document, module call and tensor",
|
|
318
|
+
"ranking": "largest absolute value first; non-finite elements are not "
|
|
319
|
+
"ranked and stay counted in the observation row",
|
|
320
|
+
"ties": "equal absolute values are ordered by ascending flat index, and the "
|
|
321
|
+
"same index order decides which of them the k-th rank keeps",
|
|
322
|
+
"position": "column of the current model input, right padding already "
|
|
323
|
+
"excluded, so it is the same coordinate before and after "
|
|
324
|
+
"removal; a decode step's own position is its generation step",
|
|
325
|
+
"layouts": "bsh and flattened_bs only; attention matrices are measured but "
|
|
326
|
+
"never unfolded into query/key positions",
|
|
327
|
+
"join": "forward, call, batch_row, module, io, tensor_path give the "
|
|
328
|
+
"observation row with finite_count and the non-finite counts"})
|
|
329
|
+
self._meta("axis_statistics", {
|
|
330
|
+
"axes": self.axes_spec,
|
|
331
|
+
"modules": config.get("axes_modules"),
|
|
332
|
+
"unit": "one row per axis index, per document, module call and tensor",
|
|
333
|
+
"feature": "reduces the document's valid positions and keeps the tensor's "
|
|
334
|
+
"own channel index",
|
|
335
|
+
"position": "reduces the feature axis and keeps the column of the current "
|
|
336
|
+
"model input, right padding already excluded",
|
|
337
|
+
"count": "elements reduced into this index, which is the other axis's length; "
|
|
338
|
+
"mean/M2/min/max describe the finite ones and the rest are counted",
|
|
339
|
+
"grid": "only the chosen axis is stored; the token x feature grid is never "
|
|
340
|
+
"expanded, which is the point of choosing an axis",
|
|
341
|
+
"alignment": "the same position index in two documents is the same column of "
|
|
342
|
+
"the model input, not a claim that those tokens mean the same "
|
|
343
|
+
"thing; a decode step is column 0, so `step` separates the steps",
|
|
344
|
+
"documents": "axis_dataset.documents counts the documents that reached this "
|
|
345
|
+
"index at all, so a position only long documents have is visibly "
|
|
346
|
+
"pooled over fewer of them",
|
|
347
|
+
"axis_role": "input_position for every position row; for a feature row, "
|
|
348
|
+
"expert at a registered router's logits, router_rank at its "
|
|
349
|
+
"selected weights or indices, and channel otherwise",
|
|
350
|
+
"layouts": "bsh and flattened_bs only; attention matrices are measured but "
|
|
351
|
+
"never unfolded into query/key positions",
|
|
352
|
+
"join": "forward, call, batch_row, module, io, tensor_path give the pooled "
|
|
353
|
+
"observation row these reduce, with its shape and source_shape"})
|
|
354
|
+
self._meta("routing", {
|
|
355
|
+
"modules": config.get("routing"),
|
|
356
|
+
"registered": sorted(ROUTER_BOUNDARIES),
|
|
357
|
+
"source": "observed - the weights and expert ids the registered router "
|
|
358
|
+
"returned, which are the tensors its parent block passes to the "
|
|
359
|
+
"experts module unchanged in the installed implementation",
|
|
360
|
+
"derived": "not stored: a top-k recomputed from logits would need the "
|
|
361
|
+
"implementation's softmax dtype, its normalisation and a k, none "
|
|
362
|
+
"of which is observable at the boundary. A router that returns "
|
|
363
|
+
"logits alone - transformers 4.57 takes the top-k in the parent "
|
|
364
|
+
"block - records routing_unavailable instead",
|
|
365
|
+
"rank": "the column of the router's own top-k output; torch.topk returns "
|
|
366
|
+
"descending probability in both installed implementations",
|
|
367
|
+
"weight_scope": "selected_top_k - the weight of one selected expert among the "
|
|
368
|
+
"k selected for that token. Whether they are normalised over "
|
|
369
|
+
"that k is the implementation's choice (Mixtral always; Qwen3 "
|
|
370
|
+
"when config.norm_topk_prob), so sum a token's ranks to see",
|
|
371
|
+
"weight_state": "finite, nan, posinf or neginf; a non-finite weight is stored "
|
|
372
|
+
"as its state with a null value, never as a number",
|
|
373
|
+
"coverage": "modules.routing says the module was selected, not that rows "
|
|
374
|
+
"exist; a selected router with none has its reason in exclusions, "
|
|
375
|
+
"and router_experts/router_top_k are filled from the first row",
|
|
376
|
+
"expert_count": "coverage.router_experts, the width of the router's logits, so "
|
|
377
|
+
"an expert with no rows is distinguishable from an index that "
|
|
378
|
+
"does not exist. Selection frequency is a count over these "
|
|
379
|
+
"rows and is not stored separately",
|
|
380
|
+
"position": "as in extremes: the column of the current model input, right "
|
|
381
|
+
"padding already excluded"})
|
|
382
|
+
self._meta("population", "all valid current input positions; float tensors only")
|
|
383
|
+
self._meta("reduction", {
|
|
384
|
+
"chunk_elements": self.chunk_elements,
|
|
385
|
+
"default": CHUNK_ELEMENTS,
|
|
386
|
+
"moments": "flattened sample split into chunks of this many elements",
|
|
387
|
+
"axis_statistics": "blocks of max(1, chunk_elements // reduced-axis length) "
|
|
388
|
+
"kept indices",
|
|
389
|
+
"invariance": "a setting of the reduction, not of the data: stored values "
|
|
390
|
+
"must not depend on it",
|
|
391
|
+
"memory": "bounds the float64 buffers only; flattening a strided sample can "
|
|
392
|
+
"still copy the whole sample"})
|
|
393
|
+
self.db.commit()
|
|
394
|
+
|
|
395
|
+
def _meta(self, key: str, value: Any) -> None:
|
|
396
|
+
self.db.execute("INSERT OR REPLACE INTO metadata VALUES (?, ?)",
|
|
397
|
+
(key, json.dumps(value)))
|
|
398
|
+
|
|
399
|
+
def resolve_modules(self, targets: Iterable[tuple[str, str, str | None]]) -> None:
|
|
400
|
+
"""Persist the complete selection, and its layouts, before the first hook.
|
|
401
|
+
|
|
402
|
+
A flattened boundary is decided here, from the model structure, and never
|
|
403
|
+
from a tensor's shape at observation time: an expert-major child can have the
|
|
404
|
+
same first axis as a token-major one on the batch that happens to route that
|
|
405
|
+
way.
|
|
406
|
+
"""
|
|
407
|
+
rows = []
|
|
408
|
+
for path, module_type, parent_type in targets:
|
|
409
|
+
rule = f"{parent_type}.{path.rpartition('.')[2]}" if parent_type else None
|
|
410
|
+
rule = rule if rule in FLATTENED_BOUNDARIES else None
|
|
411
|
+
self.flattened[path] = rule
|
|
412
|
+
k = self.extremes_k if (self.extremes_pattern is not None
|
|
413
|
+
and self.extremes_pattern.search(path)) else 0
|
|
414
|
+
self.extremes[path] = k
|
|
415
|
+
axes = self.axes_spec if (self.axes_pattern is not None
|
|
416
|
+
and self.axes_pattern.search(path)) else None
|
|
417
|
+
self.axes[path] = axes
|
|
418
|
+
# A router is one of the flattened boundaries, so its rows already map to
|
|
419
|
+
# documents. Knowing it is a router is separate from being asked for its
|
|
420
|
+
# selection: the registry also names what a feature index means there.
|
|
421
|
+
router = ROUTER_BOUNDARIES.get(rule) if rule else None
|
|
422
|
+
self.routers[path] = router
|
|
423
|
+
selected = (router is not None and self.routing_pattern is not None
|
|
424
|
+
and bool(self.routing_pattern.search(path)))
|
|
425
|
+
self.routing[path] = router if selected else None
|
|
426
|
+
rows.append((path, module_type, parent_type, rule, k or None, axes,
|
|
427
|
+
rule if router else None, "selected" if selected else None))
|
|
428
|
+
self.db.executemany("INSERT INTO modules(module, module_type, parent_type, "
|
|
429
|
+
"flattened_rule, extremes_k, axis_stats, router_rule, routing, "
|
|
430
|
+
"registration_status) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 'pending')",
|
|
431
|
+
rows)
|
|
432
|
+
self._meta("resolved_module_count", len(rows))
|
|
433
|
+
self._meta("flattened_module_count", sum(row[3] is not None for row in rows))
|
|
434
|
+
self._meta("extremes_module_count", sum(bool(row[4]) for row in rows))
|
|
435
|
+
self._meta("axis_module_count", sum(bool(row[5]) for row in rows))
|
|
436
|
+
routers = sum(bool(row[7]) for row in rows)
|
|
437
|
+
self._meta("routing_module_count", routers)
|
|
438
|
+
self.db.commit()
|
|
439
|
+
if self.routing_pattern is not None and not routers:
|
|
440
|
+
# Refuse the run rather than leave an empty table that reads like a model
|
|
441
|
+
# which never routed. This is the earliest point the module paths exist.
|
|
442
|
+
raise ValueError(
|
|
443
|
+
"--module-stats-routing matched no registered router boundary among the "
|
|
444
|
+
f"{len(rows)} selected modules. Routing selections are observable only at "
|
|
445
|
+
f"{sorted(ROUTER_BOUNDARIES)}, matched by the parent class and the "
|
|
446
|
+
"attribute holding the child; --debug-modules must select one of them")
|
|
447
|
+
|
|
448
|
+
def registration(self, module: str, status: str, error: str | None = None) -> None:
|
|
449
|
+
"""Keep partial registration distinguishable from a never-attempted target."""
|
|
450
|
+
self.db.execute("UPDATE modules SET registration_status=?, registration_error=? WHERE module=?",
|
|
451
|
+
(status, error, module))
|
|
452
|
+
self.db.commit()
|
|
453
|
+
|
|
454
|
+
def called(self, module: str) -> None:
|
|
455
|
+
"""Count one invocation before inspection, independently of its tensor count."""
|
|
456
|
+
self.db.execute("UPDATE modules SET call_count=call_count+1 WHERE module=?", (module,))
|
|
457
|
+
self.db.commit()
|
|
458
|
+
|
|
459
|
+
def set_samples(self, contexts: Iterable[Any]) -> None:
|
|
460
|
+
self.contexts = list(contexts)
|
|
461
|
+
self.next_step = 0
|
|
462
|
+
|
|
463
|
+
def begin(self, forward: int, args: tuple, kwargs: dict) -> None:
|
|
464
|
+
"""Read the real root-call shape, before any selected child module runs."""
|
|
465
|
+
import torch
|
|
466
|
+
|
|
467
|
+
self.forward = forward
|
|
468
|
+
self.lengths = None
|
|
469
|
+
self.width = self.batch = 0
|
|
470
|
+
self.stage = "unknown"
|
|
471
|
+
self.db.execute("INSERT INTO forwards VALUES (?, 'incomplete')", (forward,))
|
|
472
|
+
ids = kwargs.get("input_ids")
|
|
473
|
+
if ids is None:
|
|
474
|
+
ids = kwargs.get("inputs_embeds")
|
|
475
|
+
if ids is None and args:
|
|
476
|
+
ids = args[0]
|
|
477
|
+
if not isinstance(ids, torch.Tensor) or ids.ndim < 2 or not self.contexts:
|
|
478
|
+
return
|
|
479
|
+
self.batch, self.width = ids.shape[:2]
|
|
480
|
+
kind = getattr(self.contexts[0], "task_kind", "generate")
|
|
481
|
+
incremental = kind == "generate" and not hasattr(self.contexts[0], "positions")
|
|
482
|
+
self.stage = ("prefill" if self.next_step == 0 else "decode") if incremental else (
|
|
483
|
+
"teacher_forced" if kind == "generate" else "loglikelihood")
|
|
484
|
+
self.step = self.next_step if incremental else 0
|
|
485
|
+
self.next_step += int(incremental)
|
|
486
|
+
if incremental:
|
|
487
|
+
# The backend generates one unpadded document at a time. Recomputed
|
|
488
|
+
# prefixes (use_cache=False) are excluded rather than counted again.
|
|
489
|
+
if self.batch == 1 and (self.step == 0 or self.width == 1):
|
|
490
|
+
self.lengths = [self.width]
|
|
491
|
+
else:
|
|
492
|
+
lengths = {}
|
|
493
|
+
for ctx in self.contexts:
|
|
494
|
+
length = getattr(ctx, "input_length", None)
|
|
495
|
+
if length is not None:
|
|
496
|
+
lengths[ctx.batch_row] = int(length)
|
|
497
|
+
if set(lengths) == set(range(self.batch)):
|
|
498
|
+
proposed = [lengths[i] for i in range(self.batch)]
|
|
499
|
+
if all(0 < n <= self.width for n in proposed):
|
|
500
|
+
self.lengths = proposed
|
|
501
|
+
|
|
502
|
+
def finish_forward(self) -> None:
|
|
503
|
+
self.db.execute("UPDATE forwards SET status='complete' WHERE forward=?", (self.forward,))
|
|
504
|
+
self.db.commit()
|
|
505
|
+
|
|
506
|
+
def exclude(self, module: str, io: str, path: str, reason: str) -> None:
|
|
507
|
+
self.db.execute("""INSERT INTO exclusions VALUES (?, ?, ?, ?, 1)
|
|
508
|
+
ON CONFLICT(module, io, tensor_path, reason) DO UPDATE SET count=count+1""",
|
|
509
|
+
(module, io, path, reason))
|
|
510
|
+
|
|
511
|
+
def _documents(self) -> list[tuple[int, str, int, str]]:
|
|
512
|
+
"""The documents of this forward, one entry per physical row they share.
|
|
513
|
+
|
|
514
|
+
Choice aliases of one document collapse into a single entry so that a shared
|
|
515
|
+
row is one observation, and the aliases are preserved beside it instead.
|
|
516
|
+
"""
|
|
517
|
+
grouped: dict[tuple, set] = {}
|
|
518
|
+
for ctx in self.contexts:
|
|
519
|
+
key = (ctx.batch_row, ctx.task_name, ctx.doc_id)
|
|
520
|
+
grouped.setdefault(key, set()).add(ctx.choice_idx)
|
|
521
|
+
documents = []
|
|
522
|
+
for (row, task, doc), choices in grouped.items():
|
|
523
|
+
if not 0 <= row < self.batch:
|
|
524
|
+
raise ValueError("sample batch row is outside the model input")
|
|
525
|
+
documents.append((row, task, doc, json.dumps(sorted(choices))))
|
|
526
|
+
return documents
|
|
527
|
+
|
|
528
|
+
def _feature_role(self, module: str, path: str) -> str:
|
|
529
|
+
"""What a feature index means: an expert, a selection rank, or a channel."""
|
|
530
|
+
rule = self.routers.get(module)
|
|
531
|
+
if rule:
|
|
532
|
+
if path in rule["logits"]:
|
|
533
|
+
return "expert"
|
|
534
|
+
if path in (rule["weights"], rule["experts"]):
|
|
535
|
+
return "router_rank"
|
|
536
|
+
return "channel"
|
|
537
|
+
|
|
538
|
+
def observe(self, module: str, io: str, value: Any, call: int, phase: str,
|
|
539
|
+
module_type: str = "") -> None:
|
|
540
|
+
"""Split confirmed batch/sequence layouts, exclude padding, then reduce.
|
|
541
|
+
|
|
542
|
+
Choice aliases sharing the same physical row produce one observation per
|
|
543
|
+
document. The choices column preserves the aliases without multiplying the
|
|
544
|
+
population. Distinct choice inputs remain separate observations of a document.
|
|
545
|
+
"""
|
|
546
|
+
for path, tensor in _tensors(value, io):
|
|
547
|
+
layout, reason = self._classify_tensor(module, path, tensor, phase, module_type)
|
|
548
|
+
if reason:
|
|
549
|
+
self.exclude(module, io, path, reason)
|
|
550
|
+
continue
|
|
551
|
+
|
|
552
|
+
for row, task, doc, choices in self._documents():
|
|
553
|
+
sample = self._document_tensor(tensor, layout, row)
|
|
554
|
+
if sample is None:
|
|
555
|
+
self.exclude(module, io, path, "unsupported_key_length")
|
|
556
|
+
continue
|
|
557
|
+
# 원시 관측·극값·축별 통계가 공유하는 식별 필드다. 전체 SQL 행을
|
|
558
|
+
# 만든 뒤 앞부분을 잘라 쓰지 않고, 같은 키를 명시적으로 전달한다.
|
|
559
|
+
observation_key = [self.forward, call, row, task, doc, choices,
|
|
560
|
+
module, io, path, str(tensor.dtype), self.stage,
|
|
561
|
+
self.step, layout]
|
|
562
|
+
self._record_observation(sample, tensor.shape, observation_key, module, io, path, layout)
|
|
563
|
+
|
|
564
|
+
# The selection is a property of the whole output, not of one of its tensors:
|
|
565
|
+
# the weights and the expert ids only mean anything as a pair.
|
|
566
|
+
rule = self.routing.get(module) if io == "output" else None
|
|
567
|
+
if rule is not None:
|
|
568
|
+
self._record_routing(module, value, call, rule, phase)
|
|
569
|
+
self.db.commit()
|
|
570
|
+
|
|
571
|
+
def _classify_tensor(self, module: str, path: str, tensor: Any, phase: str,
|
|
572
|
+
module_type: str) -> tuple[str | None, str | None]:
|
|
573
|
+
"""지원 layout 또는 제외 사유를 반환한다. 둘 중 하나만 값이 있다.
|
|
574
|
+
|
|
575
|
+
판정 순서는 제외 통계의 의미를 결정한다. 여러 조건에 걸리는 tensor도
|
|
576
|
+
기존과 같은 사유로 집계되도록 phase·mask·공유값 검사를 먼저 수행한다.
|
|
577
|
+
"""
|
|
578
|
+
if phase != "model_forward":
|
|
579
|
+
return None, "not_model_forward"
|
|
580
|
+
if "attention_mask" in path.split("."):
|
|
581
|
+
return None, "attention_mask"
|
|
582
|
+
if "RotaryEmbedding" in module_type:
|
|
583
|
+
return None, "shared_or_cached"
|
|
584
|
+
if self.lengths is None:
|
|
585
|
+
return None, "missing_sample_or_valid_length"
|
|
586
|
+
if tensor is None:
|
|
587
|
+
return None, "opaque_or_deep_container"
|
|
588
|
+
if not tensor.is_floating_point():
|
|
589
|
+
return None, "nonfloating"
|
|
590
|
+
if any(name in path.split(".") for name in (
|
|
591
|
+
"position_ids", "cache_position", "position_embeddings", "past_key_values",
|
|
592
|
+
"cos", "sin")):
|
|
593
|
+
return None, "shared_or_cached"
|
|
594
|
+
if tensor.ndim == 3 and tuple(tensor.shape[:2]) == (self.batch, self.width):
|
|
595
|
+
return "bsh", None
|
|
596
|
+
if tensor.ndim == 2 and self.flattened.get(module):
|
|
597
|
+
# 등록된 경계만 batch * width 순서를 보장한다. 크기가 달라졌다면
|
|
598
|
+
# 토큰이 제거·추가·재배열됐을 수 있어 문서에 귀속시키지 않는다.
|
|
599
|
+
if tensor.shape[0] == self.batch * self.width:
|
|
600
|
+
return "flattened_bs", None
|
|
601
|
+
return None, "flattened_size_mismatch"
|
|
602
|
+
if (tensor.ndim == 4 and tensor.shape[0] == self.batch
|
|
603
|
+
and tensor.shape[2] in (1, self.width)):
|
|
604
|
+
# 4차원이라는 사실만으로 B,H,Q,K 의미를 부여하지 않는다.
|
|
605
|
+
if "attentions" in path or "Attention" in module_type or "self_attn" in module:
|
|
606
|
+
return "bhqk", None
|
|
607
|
+
return None, "unsupported_layout"
|
|
608
|
+
|
|
609
|
+
def _document_tensor(self, tensor: Any, layout: str, row: int) -> Any:
|
|
610
|
+
"""검증된 layout에서 한 문서의 padding을 제외한다. 부족한 key 축은 None.
|
|
611
|
+
|
|
612
|
+
flattened 입력은 reshape하지 않고 해당 문서의 행만 자른다. 이렇게 해야
|
|
613
|
+
non-contiguous tensor 전체를 복사하지 않고 원래 토큰 순서를 유지한다.
|
|
614
|
+
"""
|
|
615
|
+
length = self.lengths[row]
|
|
616
|
+
if layout == "bsh":
|
|
617
|
+
return tensor[row, :length]
|
|
618
|
+
if layout == "flattened_bs":
|
|
619
|
+
start = row * self.width
|
|
620
|
+
return tensor[start:start + length]
|
|
621
|
+
# bhqk: decode의 key에는 KV cache가 포함된다. batch=1이라 key padding이 없다.
|
|
622
|
+
key_length = tensor.shape[-1] if self.stage == "decode" else length
|
|
623
|
+
if tensor.shape[-1] < key_length:
|
|
624
|
+
return None
|
|
625
|
+
return tensor[row, :, :min(length, tensor.shape[2]), :key_length]
|
|
626
|
+
|
|
627
|
+
def _record_observation(self, sample: Any, source_shape: Any, observation_key: list,
|
|
628
|
+
module: str, io: str, path: str, layout: str) -> None:
|
|
629
|
+
"""한 문서의 moments와 선택된 추가 통계를 같은 식별 키로 기록한다.
|
|
630
|
+
|
|
631
|
+
commit은 observe가 담당한다. 중간에 실패했을 때 남는 관측 범위와
|
|
632
|
+
forward 완료 여부에 따른 집계 규칙을 기존과 동일하게 유지한다.
|
|
633
|
+
"""
|
|
634
|
+
stats = tensor_moments(sample, self.chunk_elements)
|
|
635
|
+
if stats["finite_count"] and not all(math.isfinite(stats[k]) for k in ("mean", "m2")):
|
|
636
|
+
# float64도 넘칠 수 있다. NaN이 SQL NULL로 저장돼 정상 평균처럼
|
|
637
|
+
# 취급되지 않도록 전체 관측을 제외한다.
|
|
638
|
+
self.exclude(module, io, path, "moment_overflow")
|
|
639
|
+
return
|
|
640
|
+
values = observation_key + [json.dumps(list(sample.shape)), json.dumps(list(source_shape))]
|
|
641
|
+
values += [stats[k] for k in ("count", "finite_count", "nan", "posinf", "neginf",
|
|
642
|
+
"mean", "m2", "min", "max")]
|
|
643
|
+
placeholders = ",".join("?" for _ in values)
|
|
644
|
+
self.db.execute(f"INSERT INTO observations VALUES ({placeholders})", values)
|
|
645
|
+
k = self.extremes.get(module, 0)
|
|
646
|
+
if k and layout in ("bsh", "flattened_bs"):
|
|
647
|
+
self._record_extremes(sample, min(k, stats["finite_count"]), observation_key)
|
|
648
|
+
if self.axes.get(module) and layout in ("bsh", "flattened_bs"):
|
|
649
|
+
if self._record_axis_stats(sample, self.axes[module], observation_key,
|
|
650
|
+
self._feature_role(module, path)):
|
|
651
|
+
self.exclude(module, io, path, "axis_moment_overflow")
|
|
652
|
+
|
|
653
|
+
def _record_extremes(self, sample: Any, k: int, keys: list) -> None:
|
|
654
|
+
"""Rank one document's slice by absolute value, and keep where the k largest are.
|
|
655
|
+
|
|
656
|
+
Only confirmed [position, feature] slices reach this, so a row's position is a
|
|
657
|
+
token of this model input and its feature is a channel of that tensor. Non-finite
|
|
658
|
+
elements are never ranked - an Inf would otherwise be every rank - and stay
|
|
659
|
+
counted in the observation row this shares a key with. Equal magnitudes are
|
|
660
|
+
ordered by ascending flat index, which also decides which of them the last rank
|
|
661
|
+
keeps, so the same tensor always writes the same rows.
|
|
662
|
+
"""
|
|
663
|
+
import torch
|
|
664
|
+
|
|
665
|
+
if k <= 0:
|
|
666
|
+
return
|
|
667
|
+
magnitude = sample.abs() # a fresh contiguous tensor
|
|
668
|
+
magnitude.masked_fill_(~torch.isfinite(sample), -1.0)
|
|
669
|
+
flat = magnitude.reshape(-1)
|
|
670
|
+
_, selected = torch.topk(flat, k)
|
|
671
|
+
threshold = flat[selected[-1]]
|
|
672
|
+
if int((flat == threshold).sum()) > int((flat[selected] == threshold).sum()):
|
|
673
|
+
# More elements share the smallest kept magnitude than there is room for.
|
|
674
|
+
above = selected[flat[selected] > threshold]
|
|
675
|
+
tied = (flat == threshold).nonzero().flatten()
|
|
676
|
+
selected = torch.cat([above, tied[:k - above.numel()]])
|
|
677
|
+
selected = selected.sort().values
|
|
678
|
+
selected = selected[torch.argsort(flat[selected], descending=True, stable=True)]
|
|
679
|
+
features = sample.shape[1]
|
|
680
|
+
positions = torch.div(selected, features, rounding_mode="floor")
|
|
681
|
+
columns = selected % features
|
|
682
|
+
# Advanced indexing, so a non-contiguous slice is gathered, not copied whole.
|
|
683
|
+
signed = sample[positions, columns].double().tolist()
|
|
684
|
+
rows = [keys + [int(position), int(column), value, rank]
|
|
685
|
+
for rank, (position, column, value) in enumerate(
|
|
686
|
+
zip(positions.tolist(), columns.tolist(), signed))]
|
|
687
|
+
self.db.executemany("INSERT INTO extremes VALUES (" + ",".join("?" * 17) + ")", rows)
|
|
688
|
+
|
|
689
|
+
def _record_axis_stats(self, sample: Any, axes: str, keys: list,
|
|
690
|
+
feature_role: str) -> bool:
|
|
691
|
+
"""Reduce one document's slice along one axis, keeping the other axis's index.
|
|
692
|
+
|
|
693
|
+
A feature row reduces this document's valid positions and keeps the tensor's own
|
|
694
|
+
channel; a position row reduces the channels and keeps the column of the current
|
|
695
|
+
model input. Only confirmed [position, feature] slices reach this, so both indices
|
|
696
|
+
mean what they are called. Non-finite elements are counted per index and excluded
|
|
697
|
+
from that index's mean/M2/min/max, exactly as the pooled observation does, so a
|
|
698
|
+
channel of NaN is visible rather than averaged into one.
|
|
699
|
+
|
|
700
|
+
Chunking is along the axis that is *kept*, never the one being reduced: each
|
|
701
|
+
block's float64 copy and its comparison masks are bounded, and no partial moment
|
|
702
|
+
has to be merged across chunks. Returns whether any index exceeded float64.
|
|
703
|
+
"""
|
|
704
|
+
import torch
|
|
705
|
+
|
|
706
|
+
positions, features = sample.shape
|
|
707
|
+
overflow = False
|
|
708
|
+
for axis in ("feature", "position"):
|
|
709
|
+
if axes not in (axis, "both"):
|
|
710
|
+
continue
|
|
711
|
+
reduced = 0 if axis == "feature" else 1
|
|
712
|
+
kept, other = ((features, positions) if axis == "feature"
|
|
713
|
+
else (positions, features))
|
|
714
|
+
# A position index is a column of the model input whatever the tensor is;
|
|
715
|
+
# only a feature index takes its meaning from the module.
|
|
716
|
+
role = feature_role if axis == "feature" else "input_position"
|
|
717
|
+
if not kept or not other:
|
|
718
|
+
continue
|
|
719
|
+
span = max(1, self.chunk_elements // other)
|
|
720
|
+
rows = []
|
|
721
|
+
for start in range(0, kept, span):
|
|
722
|
+
block = (sample[:, start:start + span] if axis == "feature"
|
|
723
|
+
else sample[start:start + span])
|
|
724
|
+
values = block.detach().to(dtype=torch.float64)
|
|
725
|
+
mask = torch.isfinite(values)
|
|
726
|
+
zero = torch.zeros((), dtype=torch.float64, device=values.device)
|
|
727
|
+
finite = mask.sum(dim=reduced)
|
|
728
|
+
mean = torch.where(mask, values, zero).sum(dim=reduced) / finite.clamp(min=1)
|
|
729
|
+
# Masked after subtracting, so an Inf does not poison its neighbours in
|
|
730
|
+
# the sum; an index whose every element is non-finite gets None below.
|
|
731
|
+
deviation = torch.where(mask, values - mean.unsqueeze(reduced), zero)
|
|
732
|
+
counters = torch.stack([finite,
|
|
733
|
+
torch.isnan(values).sum(dim=reduced),
|
|
734
|
+
torch.isposinf(values).sum(dim=reduced),
|
|
735
|
+
torch.isneginf(values).sum(dim=reduced)]).cpu().tolist()
|
|
736
|
+
moments = torch.stack([
|
|
737
|
+
mean, (deviation * deviation).sum(dim=reduced),
|
|
738
|
+
torch.where(mask, values, zero + math.inf).amin(dim=reduced),
|
|
739
|
+
torch.where(mask, values, zero - math.inf).amax(dim=reduced),
|
|
740
|
+
]).cpu().tolist()
|
|
741
|
+
for offset in range(block.shape[1 - reduced]):
|
|
742
|
+
finite_count = counters[0][offset]
|
|
743
|
+
if finite_count:
|
|
744
|
+
found = [moments[j][offset] for j in range(4)]
|
|
745
|
+
if not all(math.isfinite(value) for value in found[:2]):
|
|
746
|
+
# As in the pooled row: never store a NaN mean as SQL NULL
|
|
747
|
+
# and let it read later as a missing measurement.
|
|
748
|
+
overflow = True
|
|
749
|
+
continue
|
|
750
|
+
else:
|
|
751
|
+
found = [None, None, None, None]
|
|
752
|
+
rows.append(keys + [axis, role, start + offset, other]
|
|
753
|
+
+ [counters[j][offset] for j in range(4)] + found)
|
|
754
|
+
if rows:
|
|
755
|
+
self.db.executemany(
|
|
756
|
+
"INSERT INTO axis_stats VALUES (" + ",".join("?" * 25) + ")", rows)
|
|
757
|
+
return overflow
|
|
758
|
+
|
|
759
|
+
def _record_routing(self, module: str, value: Any, call: int, rule: dict,
|
|
760
|
+
phase: str) -> None:
|
|
761
|
+
"""Store the expert selection the router returned, per token and per rank.
|
|
762
|
+
|
|
763
|
+
This is the dispatch itself, not a reconstruction: in the installed
|
|
764
|
+
implementations the parent block passes exactly these two tensors to its experts
|
|
765
|
+
module. Rows map to documents by the flattened rule, so a padded position can
|
|
766
|
+
never appear. A router that returns logits alone leaves `routing_unavailable`
|
|
767
|
+
rather than a top-k recomputed here, because the softmax dtype, the
|
|
768
|
+
normalisation and k all belong to the parent block that took it.
|
|
769
|
+
"""
|
|
770
|
+
import torch
|
|
771
|
+
|
|
772
|
+
weights = _output_tensor(value, rule["weights"])
|
|
773
|
+
experts = _output_tensor(value, rule["experts"])
|
|
774
|
+
if phase != "model_forward":
|
|
775
|
+
reason = "not_model_forward"
|
|
776
|
+
elif self.lengths is None:
|
|
777
|
+
reason = "missing_sample_or_valid_length"
|
|
778
|
+
elif weights is None or experts is None:
|
|
779
|
+
reason = "routing_unavailable"
|
|
780
|
+
elif not weights.is_floating_point() or experts.is_floating_point():
|
|
781
|
+
reason = "routing_unexpected_dtype"
|
|
782
|
+
elif (weights.ndim != 2 or tuple(weights.shape) != tuple(experts.shape)
|
|
783
|
+
or weights.shape[0] != self.batch * self.width):
|
|
784
|
+
# The same refusal as the moments: a first axis that is not batch * width
|
|
785
|
+
# means the rows are no longer the caller's tokens, in the caller's order.
|
|
786
|
+
reason = "routing_size_mismatch"
|
|
787
|
+
else:
|
|
788
|
+
reason = None
|
|
789
|
+
if reason:
|
|
790
|
+
self.exclude(module, "output", rule["weights"], reason)
|
|
791
|
+
return
|
|
792
|
+
|
|
793
|
+
logits = next((tensor for tensor in (_output_tensor(value, spec)
|
|
794
|
+
for spec in rule["logits"])
|
|
795
|
+
if tensor is not None), None)
|
|
796
|
+
# The logits' width is the only place the expert count is observable, and it is
|
|
797
|
+
# what tells a never-selected expert apart from an index that does not exist.
|
|
798
|
+
self.db.execute("UPDATE modules SET router_experts=coalesce(?, router_experts), "
|
|
799
|
+
"router_top_k=? WHERE module=?",
|
|
800
|
+
(None if logits is None else int(logits.shape[-1]),
|
|
801
|
+
int(weights.shape[1]), module))
|
|
802
|
+
dtype = str(weights.dtype)
|
|
803
|
+
rows = []
|
|
804
|
+
for row, task, doc, choices in self._documents():
|
|
805
|
+
start = row * self.width
|
|
806
|
+
length = self.lengths[row]
|
|
807
|
+
chosen = weights[start:start + length].detach().to(
|
|
808
|
+
dtype=torch.float64).cpu().tolist()
|
|
809
|
+
ids = experts[start:start + length].detach().cpu().tolist()
|
|
810
|
+
for position, (token_weights, token_experts) in enumerate(zip(chosen, ids)):
|
|
811
|
+
for rank, (weight, expert) in enumerate(zip(token_weights, token_experts)):
|
|
812
|
+
state = ("finite" if math.isfinite(weight) else
|
|
813
|
+
"nan" if math.isnan(weight) else
|
|
814
|
+
"posinf" if weight > 0 else "neginf")
|
|
815
|
+
rows.append([self.forward, call, row, task, doc, choices, module,
|
|
816
|
+
self.stage, self.step, "flattened_bs", dtype,
|
|
817
|
+
position, rank, int(expert),
|
|
818
|
+
weight if state == "finite" else None, state,
|
|
819
|
+
rule["selection"], "selected_top_k"])
|
|
820
|
+
if rows:
|
|
821
|
+
self.db.executemany(
|
|
822
|
+
"INSERT INTO routing VALUES (" + ",".join("?" * 18) + ")", rows)
|
|
823
|
+
|
|
824
|
+
def close(self, success: bool) -> None:
|
|
825
|
+
"""Keep partial observations, but aggregate only completed forwards.
|
|
826
|
+
|
|
827
|
+
A failure to aggregate is written into this session's metadata before being
|
|
828
|
+
raised, so the session says why it has no tables instead of looking unfinished.
|
|
829
|
+
The caller decides what to do with it: `debug.py` reports it and carries on,
|
|
830
|
+
because an evaluation that has already been scored must not be lost to a failure
|
|
831
|
+
in a side artifact.
|
|
832
|
+
"""
|
|
833
|
+
if self.closed:
|
|
834
|
+
return
|
|
835
|
+
self.closed = True
|
|
836
|
+
try:
|
|
837
|
+
self._meta("status", "aggregating" if success else "failed")
|
|
838
|
+
self._meta("aggregation_status", "running")
|
|
839
|
+
self.db.commit()
|
|
840
|
+
try:
|
|
841
|
+
aggregate(self.db)
|
|
842
|
+
export_tables(self.db, self.directory)
|
|
843
|
+
except BaseException as error:
|
|
844
|
+
self._meta("status", "failed")
|
|
845
|
+
self._meta("aggregation_status", "failed")
|
|
846
|
+
self._meta("aggregation_error", f"{type(error).__name__}: {error}"[:2000])
|
|
847
|
+
self.db.commit()
|
|
848
|
+
raise
|
|
849
|
+
self._meta("status", "complete" if success else "failed")
|
|
850
|
+
self._meta("aggregation_status", "complete")
|
|
851
|
+
self.db.commit()
|
|
852
|
+
finally:
|
|
853
|
+
self.db.close()
|
|
854
|
+
|
|
855
|
+
|
|
856
|
+
def _pool_calls(db: sqlite3.Connection, source: str, target: str, group: tuple) -> None:
|
|
857
|
+
"""Pool every call of one document into one row per group key.
|
|
858
|
+
|
|
859
|
+
Only completed forwards contribute: a forward that died mid-way stays readable in
|
|
860
|
+
the raw table and never becomes part of a summary.
|
|
861
|
+
"""
|
|
862
|
+
rows = db.execute(f"SELECT o.* FROM {source} o JOIN forwards f USING (forward) "
|
|
863
|
+
"WHERE f.status='complete' ORDER BY " + ", ".join(group) + ", doc_id")
|
|
864
|
+
key = lambda row: tuple(row[k] for k in group) + (row["doc_id"],)
|
|
865
|
+
for identity, records in itertools.groupby(rows, key):
|
|
866
|
+
pooled: dict = {}
|
|
867
|
+
count = 0
|
|
868
|
+
for record in records:
|
|
869
|
+
pooled = merge_moments(pooled, dict(record))
|
|
870
|
+
count += 1
|
|
871
|
+
std = math.sqrt(pooled["m2"] / pooled["finite_count"]) if pooled["finite_count"] else None
|
|
872
|
+
values = list(identity) + [count] + [pooled[k] for k in (
|
|
873
|
+
"count", "finite_count", "nan", "posinf", "neginf", "mean", "m2", "min", "max")] + [std]
|
|
874
|
+
db.execute(f"INSERT INTO {target} VALUES (" + ",".join("?" * len(values)) + ")", values)
|
|
875
|
+
|
|
876
|
+
|
|
877
|
+
def _pool_documents(db: sqlite3.Connection, source: str, target: str, group: tuple) -> None:
|
|
878
|
+
"""Combine documents two ways: equal element weight, and equal document weight.
|
|
879
|
+
|
|
880
|
+
`documents` counts the documents present in this group at all. For a position axis
|
|
881
|
+
that is the per-position valid document count, because a position only the longer
|
|
882
|
+
documents reach is pooled over exactly those.
|
|
883
|
+
"""
|
|
884
|
+
rows = db.execute(f"SELECT * FROM {source} ORDER BY " + ", ".join(group))
|
|
885
|
+
for identity, records in itertools.groupby(rows, lambda row: tuple(row[k] for k in group)):
|
|
886
|
+
pooled, means = {}, {}
|
|
887
|
+
documents = observations = 0
|
|
888
|
+
for record in records:
|
|
889
|
+
record = dict(record)
|
|
890
|
+
documents += 1
|
|
891
|
+
observations += record["observations"]
|
|
892
|
+
pooled = merge_moments(pooled, record)
|
|
893
|
+
if record["finite_count"]:
|
|
894
|
+
means = merge_moments(means, dict(count=1, finite_count=1,
|
|
895
|
+
mean=record["mean"], m2=0.0, min=record["mean"], max=record["mean"]))
|
|
896
|
+
means = merge_moments({}, means)
|
|
897
|
+
std = math.sqrt(pooled["m2"] / pooled["finite_count"]) if pooled["finite_count"] else None
|
|
898
|
+
mean_std = math.sqrt(means["m2"] / means["finite_count"]) if means["finite_count"] else None
|
|
899
|
+
values = list(identity) + [documents, means["finite_count"], observations]
|
|
900
|
+
values += [pooled[k] for k in ("count", "finite_count", "nan", "posinf", "neginf",
|
|
901
|
+
"mean", "m2", "min", "max")]
|
|
902
|
+
values += [std, means["mean"], mean_std, means["min"], means["max"]]
|
|
903
|
+
n = means["finite_count"]
|
|
904
|
+
sample_std = math.sqrt(means["m2"] / (n - 1)) if n >= 2 else None
|
|
905
|
+
sem = sample_std / math.sqrt(n) if n >= 2 else None
|
|
906
|
+
values += [means["m2"], sample_std, sem]
|
|
907
|
+
db.execute(f"INSERT INTO {target} VALUES (" + ",".join("?" * len(values)) + ")", values)
|
|
908
|
+
|
|
909
|
+
|
|
910
|
+
def aggregate(db: sqlite3.Connection) -> None:
|
|
911
|
+
"""Pool calls per document, then compute pooled and equal-document summaries.
|
|
912
|
+
|
|
913
|
+
SQLite sorts on disk. Python holds only one group's moments at a time, so the
|
|
914
|
+
number of documents does not determine aggregation RAM. Failed/incomplete
|
|
915
|
+
forwards remain inspectable in observations and never enter these summaries.
|
|
916
|
+
|
|
917
|
+
The per-axis tables are the same two stages over the same moments, with the axis and
|
|
918
|
+
its index added to the group key, so a channel or a position is summarised exactly
|
|
919
|
+
the way the whole tensor is - and never by averaging the documents' variances.
|
|
920
|
+
"""
|
|
921
|
+
db.row_factory = sqlite3.Row
|
|
922
|
+
db.executescript("""
|
|
923
|
+
DROP TABLE IF EXISTS samples;
|
|
924
|
+
DROP TABLE IF EXISTS dataset;
|
|
925
|
+
DROP TABLE IF EXISTS axis_samples;
|
|
926
|
+
DROP TABLE IF EXISTS axis_dataset;
|
|
927
|
+
CREATE TABLE samples (
|
|
928
|
+
task_name TEXT, module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
929
|
+
layout TEXT, doc_id INTEGER, observations INTEGER, count INTEGER, finite_count INTEGER,
|
|
930
|
+
nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
931
|
+
mean REAL, m2 REAL, min REAL, max REAL, std REAL);
|
|
932
|
+
CREATE TABLE dataset (
|
|
933
|
+
task_name TEXT, module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
934
|
+
layout TEXT, documents INTEGER, finite_documents INTEGER, observations INTEGER,
|
|
935
|
+
count INTEGER, finite_count INTEGER, nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
936
|
+
pooled_mean REAL, pooled_m2 REAL, min REAL, max REAL, pooled_std REAL,
|
|
937
|
+
sample_mean REAL, sample_mean_std REAL, sample_mean_min REAL, sample_mean_max REAL,
|
|
938
|
+
sample_mean_m2 REAL, sample_mean_sample_std REAL, sample_mean_sem REAL);
|
|
939
|
+
CREATE TABLE axis_samples (
|
|
940
|
+
task_name TEXT, module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
941
|
+
layout TEXT, axis TEXT, axis_role TEXT, axis_index INTEGER,
|
|
942
|
+
doc_id INTEGER, observations INTEGER, count INTEGER, finite_count INTEGER,
|
|
943
|
+
nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
944
|
+
mean REAL, m2 REAL, min REAL, max REAL, std REAL);
|
|
945
|
+
CREATE TABLE axis_dataset (
|
|
946
|
+
task_name TEXT, module TEXT, io TEXT, tensor_path TEXT, dtype TEXT, stage TEXT,
|
|
947
|
+
layout TEXT, axis TEXT, axis_role TEXT, axis_index INTEGER,
|
|
948
|
+
documents INTEGER, finite_documents INTEGER, observations INTEGER,
|
|
949
|
+
count INTEGER, finite_count INTEGER, nan INTEGER, posinf INTEGER, neginf INTEGER,
|
|
950
|
+
pooled_mean REAL, pooled_m2 REAL, min REAL, max REAL, pooled_std REAL,
|
|
951
|
+
sample_mean REAL, sample_mean_std REAL, sample_mean_min REAL, sample_mean_max REAL,
|
|
952
|
+
sample_mean_m2 REAL, sample_mean_sample_std REAL, sample_mean_sem REAL);
|
|
953
|
+
""")
|
|
954
|
+
_pool_calls(db, "observations", "samples", GROUP)
|
|
955
|
+
_pool_documents(db, "samples", "dataset", GROUP)
|
|
956
|
+
if db.execute("SELECT 1 FROM sqlite_master WHERE name='axis_stats'").fetchone():
|
|
957
|
+
_pool_calls(db, "axis_stats", "axis_samples", AXIS_GROUP)
|
|
958
|
+
_pool_documents(db, "axis_samples", "axis_dataset", AXIS_GROUP)
|
|
959
|
+
db.commit()
|
|
960
|
+
|
|
961
|
+
|
|
962
|
+
def export_tables(db: sqlite3.Connection, directory: Path) -> None:
|
|
963
|
+
"""Export analysis tables in batches, including an explicit schema for NULLs."""
|
|
964
|
+
import pyarrow as pa
|
|
965
|
+
import pyarrow.parquet as pq
|
|
966
|
+
|
|
967
|
+
for table in ("samples", "dataset", "axis_samples", "axis_dataset", "coverage"):
|
|
968
|
+
if not db.execute("SELECT 1 FROM sqlite_master WHERE name=?", (table,)).fetchone():
|
|
969
|
+
continue
|
|
970
|
+
columns = list(db.execute(f"PRAGMA table_info({table})"))
|
|
971
|
+
schema = pa.schema([(c[1], {"TEXT": pa.string(), "INTEGER": pa.int64(),
|
|
972
|
+
"REAL": pa.float64(), "": pa.int64()}[c[2]]) for c in columns])
|
|
973
|
+
cursor = db.execute(f"SELECT * FROM {table}")
|
|
974
|
+
temporary = directory / f".{table}.parquet.tmp"
|
|
975
|
+
with pq.ParquetWriter(temporary, schema) as writer:
|
|
976
|
+
while rows := cursor.fetchmany(2048):
|
|
977
|
+
writer.write_table(pa.Table.from_pylist([dict(r) for r in rows], schema=schema))
|
|
978
|
+
os.replace(temporary, directory / f"{table}.parquet")
|
|
979
|
+
|
|
980
|
+
|
|
981
|
+
def resolve_database(path: str, pass_name: str = "trace") -> Path:
|
|
982
|
+
"""A direct database/session path, or the latest session of the selected pass."""
|
|
983
|
+
given = Path(path)
|
|
984
|
+
if given.is_file():
|
|
985
|
+
return given
|
|
986
|
+
if (given / "statistics.sqlite").is_file():
|
|
987
|
+
return given / "statistics.sqlite"
|
|
988
|
+
found = sorted((given / "debug" / "module_stats" / pass_name).glob("*/statistics.sqlite"))
|
|
989
|
+
if not found:
|
|
990
|
+
raise FileNotFoundError(f"no module statistics under {path}; run with --module-stats")
|
|
991
|
+
return found[-1]
|
|
992
|
+
|
|
993
|
+
|
|
994
|
+
def read_statistics(path: str, table: str = "dataset", pass_name: str = "trace"):
|
|
995
|
+
"""Load one analysis table as a DataFrame; never combine runs implicitly."""
|
|
996
|
+
import pandas as pd
|
|
997
|
+
|
|
998
|
+
if table not in {"dataset", "samples", "observations", "exclusions", "forwards",
|
|
999
|
+
"coverage", "metadata", "extremes", "axis_stats", "axis_samples",
|
|
1000
|
+
"axis_dataset", "routing"}:
|
|
1001
|
+
raise ValueError(f"unknown module statistics table: {table}")
|
|
1002
|
+
database = resolve_database(path, pass_name)
|
|
1003
|
+
with sqlite3.connect(database.resolve().as_uri() + "?mode=ro", uri=True) as db:
|
|
1004
|
+
if not db.execute("SELECT 1 FROM sqlite_master WHERE type IN ('table', 'view') AND name=?", (table,)).fetchone():
|
|
1005
|
+
if table in ("coverage", "extremes", "axis_stats", "axis_samples",
|
|
1006
|
+
"axis_dataset", "routing"):
|
|
1007
|
+
raise RuntimeError(f"module {table} is unavailable in this historical schema")
|
|
1008
|
+
raise RuntimeError("module statistics have not been aggregated yet; the session is "
|
|
1009
|
+
"still running or was interrupted. Completed observations remain "
|
|
1010
|
+
"available in statistics.sqlite")
|
|
1011
|
+
return pd.read_sql_query(f"SELECT * FROM {table}", db)
|
|
1012
|
+
|
|
1013
|
+
|
|
1014
|
+
def report_statistics(path: str, pass_name: str = "trace", doc: str | None = None,
|
|
1015
|
+
module: str | None = None) -> str:
|
|
1016
|
+
"""Render an explicitly scoped sample or dataset summary and its coverage."""
|
|
1017
|
+
import re
|
|
1018
|
+
|
|
1019
|
+
database = resolve_database(path, pass_name)
|
|
1020
|
+
with sqlite3.connect(database.resolve().as_uri() + "?mode=ro", uri=True) as db:
|
|
1021
|
+
status = json.loads(db.execute("SELECT value FROM metadata WHERE key='status'").fetchone()[0])
|
|
1022
|
+
aggregation_status = json.loads(db.execute(
|
|
1023
|
+
"SELECT value FROM metadata WHERE key='aggregation_status'").fetchone()[0])
|
|
1024
|
+
failure = db.execute("SELECT value FROM metadata WHERE key='aggregation_error'").fetchone()
|
|
1025
|
+
if failure:
|
|
1026
|
+
aggregation_status += f" ({json.loads(failure[0])})"
|
|
1027
|
+
incomplete = db.execute("SELECT count(*) FROM forwards WHERE status!='complete'").fetchone()[0]
|
|
1028
|
+
has_coverage = db.execute("SELECT 1 FROM sqlite_master WHERE name='coverage'").fetchone()
|
|
1029
|
+
counted = {}
|
|
1030
|
+
for table in ("extremes", "axis_stats", "routing"):
|
|
1031
|
+
counted[table] = (db.execute(f"SELECT count(*) FROM {table}").fetchone()[0]
|
|
1032
|
+
if db.execute("SELECT 1 FROM sqlite_master WHERE name=?",
|
|
1033
|
+
(table,)).fetchone() else None)
|
|
1034
|
+
extremes = counted["extremes"]
|
|
1035
|
+
module_coverage = read_statistics(str(database), "coverage") if has_coverage else None
|
|
1036
|
+
if module_coverage is not None and module:
|
|
1037
|
+
module_coverage = module_coverage[module_coverage.module.astype(str).str.contains(module, regex=True)]
|
|
1038
|
+
coverage_text = (module_coverage.to_string(index=False) if module_coverage is not None
|
|
1039
|
+
else "unavailable in this historical schema")
|
|
1040
|
+
try:
|
|
1041
|
+
frame = read_statistics(str(database), "samples" if doc is not None else "dataset")
|
|
1042
|
+
except RuntimeError:
|
|
1043
|
+
# Abrupt termination can leave durable coverage before aggregation exists.
|
|
1044
|
+
# Report that state without fabricating an empty completed dataset.
|
|
1045
|
+
if aggregation_status == "complete":
|
|
1046
|
+
raise
|
|
1047
|
+
return (f"statistics: {database}\nstatus: {status}; aggregation: {aggregation_status}; "
|
|
1048
|
+
f"incomplete forwards excluded: {incomplete}\n"
|
|
1049
|
+
"Aggregate tables unavailable; last committed module coverage "
|
|
1050
|
+
"(session-wide, all phases, includes incomplete forwards):\n" + coverage_text)
|
|
1051
|
+
if module:
|
|
1052
|
+
pattern = re.compile(module)
|
|
1053
|
+
frame = frame[frame.module.map(lambda name: bool(pattern.search(name))).astype(bool)]
|
|
1054
|
+
if doc is not None:
|
|
1055
|
+
task, _, number = doc.rpartition("#")
|
|
1056
|
+
frame = frame[frame.doc_id == int(number)]
|
|
1057
|
+
if task:
|
|
1058
|
+
frame = frame[frame.task_name == task]
|
|
1059
|
+
excluded = read_statistics(str(database), "exclusions")
|
|
1060
|
+
columns = ["task_name", "module", "io", "tensor_path", "stage"]
|
|
1061
|
+
columns += ["layout"] if "layout" in frame.columns else []
|
|
1062
|
+
columns += (["doc_id", "count", "finite_count", "mean", "std", "min", "max", "nan", "posinf", "neginf"]
|
|
1063
|
+
if doc is not None else ["documents", "finite_count", "sample_mean", "sample_mean_std",
|
|
1064
|
+
"pooled_mean", "pooled_std", "nan", "posinf", "neginf"])
|
|
1065
|
+
if doc is None:
|
|
1066
|
+
columns += [name for name in ("finite_documents", "sample_mean_m2",
|
|
1067
|
+
"sample_mean_sample_std", "sample_mean_sem") if name in frame.columns]
|
|
1068
|
+
coverage = (excluded.groupby("reason")["count"].sum().to_string()
|
|
1069
|
+
if not excluded.empty else "none")
|
|
1070
|
+
return (f"statistics: {database}\nstatus: {status}; aggregation: {aggregation_status}; "
|
|
1071
|
+
f"incomplete forwards excluded: {incomplete}\n"
|
|
1072
|
+
"sample_mean: equal document weights; pooled_mean: equal finite-element weights.\n"
|
|
1073
|
+
"Population: observed valid positions, separated by task/module/tensor/stage.\n"
|
|
1074
|
+
"sample_mean_std uses n; sample_mean_sample_std uses n-1; SEM assumes independent documents.\n"
|
|
1075
|
+
+ frame[columns].to_string(index=False) + "\nExclusions (events; units in metadata):\n" + coverage
|
|
1076
|
+
+ "\nModule coverage (session-wide, all phases, includes incomplete forwards):\n" + coverage_text
|
|
1077
|
+
+ (f"\nExtreme positions: {extremes} rows; read with "
|
|
1078
|
+
"read_statistics(path, 'extremes')." if extremes else "")
|
|
1079
|
+
+ (f"\nPer-axis statistics: {counted['axis_stats']} rows; read with "
|
|
1080
|
+
"read_statistics(path, 'axis_dataset') or 'axis_samples'."
|
|
1081
|
+
if counted["axis_stats"] else "")
|
|
1082
|
+
+ (f"\nRouting selections: {counted['routing']} rows; read with "
|
|
1083
|
+
"read_statistics(path, 'routing'); expert counts in coverage."
|
|
1084
|
+
if counted["routing"] else ""))
|