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.
@@ -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 ""))