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/storage.py ADDED
@@ -0,0 +1,1546 @@
1
+ """Schema declaration and on-disk layout for a run directory.
2
+
3
+ This module is the single source of truth for what a run writes to disk.
4
+ The schema is declared *in code* rather than in a hand-maintained markdown file, so that documentation cannot drift away from the parquet files that are actually produced.
5
+
6
+ Everything a reader needs is therefore derived from `TABLES` below:
7
+
8
+ * `ShardWriter` builds the pyarrow schema from it,
9
+ * the column descriptions are embedded into every parquet file as key/value metadata (so a stray file is self-describing),
10
+ * `describe_schema()` prints it for humans,
11
+ * `report.py` reads `layer_domain` / `comparable_across_vocab` from it to decide what may be plotted on the same axis.
12
+
13
+ Directory layout produced by a run:
14
+
15
+ results.json scores + manifest + list of signal files
16
+ samples.jsonl sample results and scored inputs
17
+ docs/part-0000.parquet document/choice grading
18
+ steps/part-0000.parquet recorded positions and token IDs
19
+ signals/part-0000.parquet logit-lens signals
20
+ similarity/part-0000.parquet layer-pair similarities
21
+ attn_norm/part-0000.parquet opt-in (--save-attention)
22
+ attention/*.safetensors opt-in (--save-attention)
23
+ raw/*.safetensors opt-in (--save-hidden)
24
+ custom/<name>/ custom hook schemas, rows and optional tensors
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import json
30
+ import os
31
+ from dataclasses import dataclass
32
+ from datetime import datetime, timezone
33
+ from typing import Any, Iterable, Literal, Sequence
34
+
35
+ import pyarrow as pa
36
+ import pyarrow.parquet as pq
37
+
38
+
39
+ # -------------------------------------------------------------------------- Fixed settings
40
+ #
41
+ # These are deliberately *not* CLI arguments: there is no reason to choose a different value, and exposing them would multiply the number of run configurations we have to reason about. They are still written into the manifest, so that a run made today stays interpretable after a constant here is changed tomorrow.
42
+ # --------------------------------------------------------------------------
43
+
44
+ # : Version of the on-disk column layout.
45
+ # Bump when columns are added/removed : or retyped.
46
+ # `report.py` refuses to mix runs with different values.
47
+ SCHEMA_VERSION = "0.4"
48
+
49
+ #: The only supported lm-eval backend name (registered in `backend.py`).
50
+ BACKEND_NAME = "hf-traced"
51
+
52
+ # : Start a new parquet shard every this many documents.
53
+ # Shards are never : appended to, so a crash leaves every already-written shard intact.
54
+ SHARD_SIZE = 500
55
+
56
+ # : Seed for picking which documents get collectioned (--save-attention / : --save-hidden).
57
+ # Fixed, so the same run always selects the same documents.
58
+ SAMPLING_SEED = 1234
59
+
60
+ #: dtype used when dumping raw hidden states to safetensors.
61
+ RAW_DTYPE = "fp16"
62
+
63
+ # : report x-axis convention.
64
+ # Layers are plotted at relative depth l/L, never : at their absolute index, so models with different depths can be overlaid.
65
+ USE_RELATIVE_DEPTH = True
66
+
67
+ # : Attention sinks (the mass that piles up on the first token) are kept as-is; : the report does not filter them out.
68
+ KEEP_ATTENTION_SINK = True
69
+
70
+ FIXED_SETTINGS: dict[str, Any] = {
71
+ "backend_name": BACKEND_NAME,
72
+ "shard_size": SHARD_SIZE,
73
+ "sampling_seed": SAMPLING_SEED,
74
+ "raw_dtype": RAW_DTYPE,
75
+ "relative_depth_axis": USE_RELATIVE_DEPTH,
76
+ "keep_attention_sink": KEEP_ATTENTION_SINK,
77
+ }
78
+
79
+
80
+ # --------------------------------------------------------------------------
81
+ # Schema declaration
82
+ # --------------------------------------------------------------------------
83
+
84
+ # : Which layer index axis a table lives on. : : - "residual": index column is `layer`, range 0..L.
85
+ # `layer j` is the *input* : of block j; `layer L` is the output of the last block. : - "block": index column is `block`, range 0..L-1.
86
+ # `block j` is the : attention *inside* block j. : - None: the table has no layer axis at all (docs, steps).
87
+ LayerDomain = Literal["residual", "block"]
88
+
89
+
90
+ @dataclass(frozen=True)
91
+ class ColumnSpec:
92
+ """One column of one table.
93
+
94
+ Attributes:
95
+ name: Column name as it appears in the parquet file.
96
+ dtype: pyarrow type.
97
+ Declared here so the writer never has to guess a type from the first batch of rows it happens to see.
98
+ description: Human readable meaning; embedded into the parquet key/value metadata.
99
+ unit: Physical unit or scale ("probability", "count", ...).
100
+ "" when the column is an identifier or a flag.
101
+ comparable_across_vocab: Whether values may be compared between models with different vocabularies.
102
+ `cos` is a cosine and therefore comparable; `target_rank` is a raw rank out of |V| and therefore is not, while `target_percentile` divides it out and is.
103
+ """
104
+
105
+ name: str
106
+ dtype: pa.DataType
107
+ description: str
108
+ unit: str = ""
109
+ comparable_across_vocab: bool = True
110
+
111
+
112
+ @dataclass(frozen=True)
113
+ class TableSpec:
114
+ """One parquet table, i.e. one subdirectory of the run directory.
115
+
116
+ Signals have different axis shapes, so they get different tables.
117
+ Forcing them into one wide table would produce rows that are mostly nulls.
118
+
119
+ Attributes:
120
+ name: Subdirectory name (`signals` -> `signals/part-0000.parquet`).
121
+ key: Columns that together identify a row.
122
+ columns: All columns, key columns first.
123
+ layer_domain: See `LayerDomain`.
124
+ description: What one row of this table means.
125
+ """
126
+
127
+ name: str
128
+ key: tuple[str, ...]
129
+ columns: tuple[ColumnSpec, ...]
130
+ layer_domain: LayerDomain | None
131
+ description: str
132
+
133
+ @property
134
+ def column_names(self) -> list[str]:
135
+ return [c.name for c in self.columns]
136
+
137
+ def arrow_schema(self) -> pa.Schema:
138
+ """Build the pyarrow schema, with column docs attached as metadata.
139
+
140
+ Example:
141
+ >>> schema = TABLES["similarity"].arrow_schema()
142
+ >>> schema.field("cos").type
143
+ DataType(float)
144
+ >>> json.loads(schema.metadata[b"eval_framework"])["layer_domain"]
145
+ 'residual'
146
+ """
147
+ fields = [pa.field(c.name, c.dtype) for c in self.columns]
148
+ meta = {
149
+ "schema_version": SCHEMA_VERSION,
150
+ "table": self.name,
151
+ "description": self.description,
152
+ "key": list(self.key),
153
+ "layer_domain": self.layer_domain,
154
+ "columns": {
155
+ c.name: {
156
+ "type": str(c.dtype),
157
+ "unit": c.unit,
158
+ "description": c.description,
159
+ "comparable_across_vocab": c.comparable_across_vocab,
160
+ }
161
+ for c in self.columns
162
+ },
163
+ }
164
+ return pa.schema(
165
+ fields, metadata={b"eval_framework": json.dumps(meta).encode("utf-8")}
166
+ )
167
+
168
+
169
+ # The four axes shared by every signal table.
170
+ # `task_name` + `doc_id` is what lm-eval uses to identify a document; `choice_idx` separates the per-choice forwards of a multiple-choice document; `step` is the position inside the scored span (loglikelihood) or the decoding step (generate).
171
+ _AXIS_TASK = ColumnSpec("task_name", pa.string(), "lm-eval task name.")
172
+ _AXIS_DOC = ColumnSpec(
173
+ "doc_id", pa.int64(), "Document id assigned by lm-eval. Unique only together with task_name."
174
+ )
175
+ _AXIS_CHOICE = ColumnSpec(
176
+ "choice_idx",
177
+ pa.int32(),
178
+ "Which choice of a multiple-choice document this forward pass was for. 0 for generate tasks.",
179
+ )
180
+ _AXIS_STEP = ColumnSpec(
181
+ "step",
182
+ pa.int32(),
183
+ "Position index. generate: decoding step 0..N. loglikelihood: position inside the scored span.",
184
+ )
185
+
186
+
187
+ TABLES: dict[str, TableSpec] = {
188
+ "docs": TableSpec(
189
+ name="docs",
190
+ key=("task_name", "doc_id", "choice_idx"),
191
+ layer_domain=None,
192
+ description=(
193
+ "One row per forward pass we ran, carrying lm-eval's grading result. "
194
+ "Grading is never recomputed here; it is joined in from --log_samples. "
195
+ "Exactly one row per (task_name, doc_id, choice_idx): a task with "
196
+ "several filters is reduced to its primary one, named in `filter`."
197
+ ),
198
+ columns=(
199
+ _AXIS_TASK,
200
+ _AXIS_DOC,
201
+ _AXIS_CHOICE,
202
+ ColumnSpec(
203
+ "is_correct",
204
+ pa.bool_(),
205
+ "Did the model answer this document correctly? A per-document value, "
206
+ "so it repeats across the choices of one document.",
207
+ ),
208
+ ColumnSpec(
209
+ "is_target_choice",
210
+ pa.bool_(),
211
+ "Is this row the forward pass of the *gold* choice? Distinguishes signals "
212
+ "taken from the correct choice from signals taken from a distractor.",
213
+ ),
214
+ ColumnSpec("predicted", pa.string(), "The answer lm-eval says the model chose."),
215
+ ColumnSpec("target", pa.string(), "The gold answer."),
216
+ ColumnSpec(
217
+ "choice_logprob",
218
+ pa.float32(),
219
+ "Summed log probability of this choice's continuation.",
220
+ unit="log probability",
221
+ comparable_across_vocab=False,
222
+ ),
223
+ ColumnSpec(
224
+ "filter",
225
+ pa.string(),
226
+ "Which lm-eval filter produced is_correct/predicted. A task can "
227
+ "define several (gsm8k has strict-match and flexible-extract) and "
228
+ "lm-eval logs one sample record per filter; keeping them all would "
229
+ "make this table non-unique on its key and fan out every join. The "
230
+ "other filters' results stay in samples.jsonl.",
231
+ ),
232
+ ),
233
+ ),
234
+ "steps": TableSpec(
235
+ name="steps",
236
+ key=("task_name", "doc_id", "choice_idx", "step"),
237
+ layer_domain=None,
238
+ description=(
239
+ "What the target token of each step was. Cannot be replaced by the layer-L row "
240
+ "of `signals`: for loglikelihood the scored token is the gold continuation token, "
241
+ "not the model's prediction, and for sampled generation the emitted token may "
242
+ "differ from the last layer's top-1."
243
+ ),
244
+ columns=(
245
+ _AXIS_TASK,
246
+ _AXIS_DOC,
247
+ _AXIS_CHOICE,
248
+ _AXIS_STEP,
249
+ ColumnSpec(
250
+ "token_id",
251
+ pa.int64(),
252
+ "generate: the token actually generated at this step. "
253
+ "loglikelihood: the continuation token being scored. "
254
+ "A generate task records every decoding step that ran, which "
255
+ "includes the tokens forming a stop sequence: generation only halts "
256
+ "after emitting it, and lm-eval strips it from the text it reports. "
257
+ "So decoding this column yields lm-eval's `resps` as a prefix, "
258
+ "usually with a stop string after it - not an exact match.",
259
+ comparable_across_vocab=False,
260
+ ),
261
+ ColumnSpec("token", pa.string(), "Tokenizer output for token_id, kept verbatim.",
262
+ comparable_across_vocab=False),
263
+ ColumnSpec(
264
+ "position",
265
+ pa.int32(),
266
+ "Absolute position in the input sequence. Prompt lengths differ per document, "
267
+ "so `step` alone does not locate the token.",
268
+ unit="tokens",
269
+ ),
270
+ ),
271
+ ),
272
+ "signals": TableSpec(
273
+ name="signals",
274
+ key=("task_name", "doc_id", "choice_idx", "step", "layer"),
275
+ layer_domain="residual",
276
+ description=(
277
+ "Logit-lens top-1 and gold-token rank, one row per residual layer per step."
278
+ ),
279
+ columns=(
280
+ _AXIS_TASK,
281
+ _AXIS_DOC,
282
+ _AXIS_CHOICE,
283
+ _AXIS_STEP,
284
+ ColumnSpec(
285
+ "layer",
286
+ pa.int32(),
287
+ "Residual index 0..L. `layer j` is the input of block j; `layer L` is the "
288
+ "output of the last block.",
289
+ ),
290
+ ColumnSpec(
291
+ "lens_token_id",
292
+ pa.int64(),
293
+ "Top-1 token id after decoding this layer through the model's own final "
294
+ "norm + lm_head. Note: on tied-embedding models `layer 0` usually decodes "
295
+ "back to the input token itself. That is a property of the convention, "
296
+ "not a bug, and must not be read as a prediction.",
297
+ comparable_across_vocab=False,
298
+ ),
299
+ ColumnSpec(
300
+ "lens_prob",
301
+ pa.float32(),
302
+ "Softmax probability of the top-1 token.",
303
+ unit="probability",
304
+ comparable_across_vocab=False,
305
+ ),
306
+ ColumnSpec(
307
+ "lens_token",
308
+ pa.string(),
309
+ "Tokenizer output for lens_token_id, stored verbatim. Leading-space markers "
310
+ "and partial byte pieces are NOT cleaned up; byte-level BPE can split a "
311
+ "character in half, so a single token may decode to a broken glyph.",
312
+ comparable_across_vocab=False,
313
+ ),
314
+ ColumnSpec(
315
+ "lens_is_special",
316
+ pa.bool_(),
317
+ "Whether lens_token_id is a special token, as judged by the tokenizer.",
318
+ comparable_across_vocab=False,
319
+ ),
320
+ ColumnSpec(
321
+ "target_rank",
322
+ pa.int64(),
323
+ "Rank of the gold token in this layer's distribution (0 = top-1). "
324
+ "loglikelihood only; null for generate tasks. Being a count rather "
325
+ "than a quantity, it is not reproducible to the last place: where "
326
+ "several tokens are nearly tied, a last-bit difference in the logits "
327
+ "moves it a place or two. Recomputing it independently on Qwen2.5-0.5B "
328
+ "moved 1.6% of rows by at most 3 out of a 151936-token vocabulary, "
329
+ "with no directional bias. Treat single-place differences as noise.",
330
+ unit="rank",
331
+ comparable_across_vocab=False,
332
+ ),
333
+ ColumnSpec(
334
+ "target_percentile",
335
+ pa.float32(),
336
+ "target_rank / vocab_size. Normalised, so it is comparable across models "
337
+ "with different vocabulary sizes.",
338
+ unit="fraction",
339
+ comparable_across_vocab=True,
340
+ ),
341
+ ),
342
+ ),
343
+ "similarity": TableSpec(
344
+ name="similarity",
345
+ key=("task_name", "doc_id", "choice_idx", "step", "layer_i", "layer_j"),
346
+ layer_domain="residual",
347
+ description=(
348
+ "Pairwise cosine similarity between residual layers at one position. "
349
+ "The matrix is symmetric, so only the upper triangle including the diagonal "
350
+ "is stored (layer_i <= layer_j)."
351
+ ),
352
+ columns=(
353
+ _AXIS_TASK,
354
+ _AXIS_DOC,
355
+ _AXIS_CHOICE,
356
+ _AXIS_STEP,
357
+ ColumnSpec("layer_i", pa.int32(), "Residual index 0..L."),
358
+ ColumnSpec("layer_j", pa.int32(), "Residual index 0..L, always >= layer_i."),
359
+ ColumnSpec(
360
+ "cos",
361
+ pa.float32(),
362
+ "Cosine similarity without centering. Interior values cluster near 1 "
363
+ "because the residual stream is only ever added to, giving every token a "
364
+ "large shared component; read relative changes between layer pairs, not "
365
+ "absolute size. The two endpoint pairs are structurally different and "
366
+ "should not be read as part of that trend: (0, 1) crosses from the "
367
+ "embedding output into the first block, and (L-1, L) crosses the final "
368
+ "norm, since layer L is the normed hidden state. Both are far lower than "
369
+ "their neighbours - on Qwen2.5-1.5B, 0.07 and 0.27 against an interior "
370
+ "0.73 to 0.98.",
371
+ unit="cosine",
372
+ comparable_across_vocab=True,
373
+ ),
374
+ ),
375
+ ),
376
+ "attn_norm": TableSpec(
377
+ name="attn_norm",
378
+ key=("task_name", "doc_id", "choice_idx", "step", "block", "head"),
379
+ layer_domain="block",
380
+ description=(
381
+ "Per-head value-vector norm. Written only with --save-attention. "
382
+ "One scalar per head per step, so unlike the attention weight tensors this "
383
+ "table is dense along the step axis."
384
+ ),
385
+ columns=(
386
+ _AXIS_TASK,
387
+ _AXIS_DOC,
388
+ _AXIS_CHOICE,
389
+ _AXIS_STEP,
390
+ ColumnSpec(
391
+ "block",
392
+ pa.int32(),
393
+ "Block index 0..L-1. `block j` is the attention inside block j: it reads "
394
+ "`layer j` and contributes to `layer j+1`.",
395
+ ),
396
+ ColumnSpec(
397
+ "head",
398
+ pa.int32(),
399
+ "Query head index. Under GQA several query heads share one key/value head, "
400
+ "so the same value norm repeats across the heads of a group.",
401
+ ),
402
+ ColumnSpec(
403
+ "value_norm",
404
+ pa.float32(),
405
+ "L2 norm of this head's value vector at the current position. "
406
+ "Captured at the projection, so it is computed from the *normalised* "
407
+ "block input - a pre-norm block applies input_layernorm before "
408
+ "attention - not from the raw residual stream. Recomputing it from "
409
+ "hidden_states without that norm is off by orders of magnitude, "
410
+ "because the residual norm grows with depth and the normalised input "
411
+ "does not.",
412
+ unit="L2 norm",
413
+ comparable_across_vocab=True,
414
+ ),
415
+ ),
416
+ ),
417
+ }
418
+
419
+
420
+ def describe_schema() -> str:
421
+ """Render the declared schema as a table for humans.
422
+
423
+ This is the documentation.
424
+ `load_signals()` prints it, and the README's output section is expected to quote this rather than restate it.
425
+
426
+ Example:
427
+ >>> print(describe_schema()) # doctest: +ELLIPSIS
428
+ schema_version 0.3
429
+ <BLANKLINE>
430
+ == docs == key=(task_name, doc_id, choice_idx) layer_domain=-
431
+ ...
432
+ """
433
+ lines = [f"schema_version {SCHEMA_VERSION}", ""]
434
+ for spec in TABLES.values():
435
+ lines.append(
436
+ f"== {spec.name} == key=({', '.join(spec.key)}) "
437
+ f"layer_domain={spec.layer_domain or '-'}"
438
+ )
439
+ lines.append(f" {spec.description}")
440
+ for col in spec.columns:
441
+ vocab = "any-vocab" if col.comparable_across_vocab else "same-vocab-only"
442
+ unit = f" [{col.unit}]" if col.unit else ""
443
+ lines.append(f" - {col.name}: {col.dtype}{unit} ({vocab})")
444
+ lines.append(f" {col.description}")
445
+ lines.append("")
446
+ return "\n".join(lines)
447
+
448
+
449
+ # --------------------------------------------------------------------------
450
+ # Writing: parquet shards
451
+ # --------------------------------------------------------------------------
452
+
453
+
454
+ class ShardWriter:
455
+ """Buffers rows for one table and flushes them into numbered parquet shards.
456
+
457
+ Shards are write-once: we never append to an existing file.
458
+ If a run dies halfway through, every shard already on disk is a complete, readable parquet file.
459
+
460
+ Numbering continues after whatever is already in the directory, so a restart adds shards rather than overwriting them.
461
+ Restarting with numbering reset to zero is worse than losing the old data outright: the shorter of the two runs only overwrites its own prefix, and what is left on disk is a mixture of both runs that still reads as a valid table.
462
+
463
+ Example:
464
+ >>> w = ShardWriter("/tmp/run", TABLES["docs"]) # doctest: +SKIP
465
+ >>> w.add({"task_name": "xnli_ko", "doc_id": 0, "choice_idx": 0,
466
+ ... "is_correct": True, "is_target_choice": True,
467
+ ... "predicted": "0", "target": "0", "choice_logprob": -1.5})
468
+ >>> w.flush() # writes /tmp/run/docs/part-0000.parquet
469
+ >>> w.close()
470
+ """
471
+
472
+ def __init__(self, run_dir: str | os.PathLike, spec: TableSpec) -> None:
473
+ self.spec = spec
474
+ self.dir = os.path.join(str(run_dir), spec.name)
475
+ self._schema = spec.arrow_schema()
476
+ self._buffer: list[dict[str, Any]] = []
477
+ self._shard_index = _next_shard_index(self.dir)
478
+
479
+ def add(self, row: dict[str, Any]) -> None:
480
+ """Buffer one row.
481
+ Missing columns become null."""
482
+ self._buffer.append(row)
483
+
484
+ def extend(self, rows: Iterable[dict[str, Any]]) -> None:
485
+ """Buffer many rows at once."""
486
+ self._buffer.extend(rows)
487
+
488
+ def flush(self) -> str | None:
489
+ """Write the buffer as the next shard.
490
+ Returns the path, or None if empty."""
491
+ if not self._buffer:
492
+ return None
493
+ os.makedirs(self.dir, exist_ok=True)
494
+ # Build columns explicitly from the declared schema so that a row that forgot a column fails loudly here rather than producing a file whose column set silently differs from every other run.
495
+ columns = {
496
+ name: [row.get(name) for row in self._buffer] for name in self.spec.column_names
497
+ }
498
+ table = pa.Table.from_pydict(columns, schema=self._schema)
499
+ path = os.path.join(self.dir, f"part-{self._shard_index:04d}.parquet")
500
+ # A process killed inside `write_table` left a footerless shard under its real name on RunPod, and resume then refused the directory.
501
+ # Write under a name no reader matches (readers take `*.parquet`), make it durable, then publish it with one rename.
502
+ import uuid
503
+ temporary = os.path.join(self.dir, f".part-{self._shard_index:04d}.parquet.{uuid.uuid4().hex}.tmp")
504
+ pq.write_table(table, temporary, compression="zstd")
505
+ with open(temporary, "rb") as stream:
506
+ os.fsync(stream.fileno())
507
+ os.replace(temporary, path)
508
+ _sync_directory(self.dir)
509
+ self._buffer.clear()
510
+ self._shard_index += 1
511
+ return path
512
+
513
+ def close(self) -> None:
514
+ """Flush whatever is left in the buffer."""
515
+ self.flush()
516
+
517
+
518
+ def _next_shard_index(directory: str) -> int:
519
+ """The first shard number not already taken in `directory`.
520
+
521
+ Example:
522
+ >>> _next_shard_index("/tmp/empty")
523
+ 0
524
+ """
525
+ if not os.path.isdir(directory):
526
+ return 0
527
+ used = [
528
+ int(name[len("part-") : -len(".parquet")])
529
+ for name in os.listdir(directory)
530
+ if name.startswith("part-") and name.endswith(".parquet")
531
+ and name[len("part-") : -len(".parquet")].isdigit()
532
+ ]
533
+ return max(used) + 1 if used else 0
534
+
535
+
536
+ class TensorStore:
537
+ """Writes opt-in raw tensors (attention weights, hidden states) as safetensors.
538
+
539
+ One file per document and choice, containing all requested layers.
540
+ Saving the same key replaces its file. This store does not provide the
541
+ transactional resume contract used by custom collection artifacts.
542
+
543
+ Example:
544
+ >>> store = TensorStore("/tmp/run", "attention") # doctest: +SKIP
545
+ >>> store.save("xnli_ko", 12, 0, {"block_00": weights},
546
+ ... meta={"step": "0"})
547
+ '/tmp/run/attention/xnli_ko__000012__0.safetensors'
548
+ """
549
+
550
+ def __init__(self, run_dir: str | os.PathLike, name: str) -> None:
551
+ self.dir = os.path.join(str(run_dir), name)
552
+ self.name = name
553
+ self._count = 0
554
+
555
+ def path_for(self, task_name: str, doc_id: int, choice_idx: int) -> str:
556
+ """Deterministic file name, so existence alone answers "did we do this?"."""
557
+ safe_task = "".join(ch if ch.isalnum() or ch in "-_." else "_" for ch in task_name)
558
+ return os.path.join(self.dir, f"{safe_task}__{doc_id:06d}__{choice_idx}.safetensors")
559
+
560
+ def save(
561
+ self,
562
+ task_name: str,
563
+ doc_id: int,
564
+ choice_idx: int,
565
+ tensors: dict[str, Any],
566
+ meta: dict[str, str] | None = None,
567
+ ) -> str:
568
+ from safetensors.torch import save_file # lazy: keeps storage.py torch-free
569
+
570
+ os.makedirs(self.dir, exist_ok=True)
571
+ path = self.path_for(task_name, doc_id, choice_idx)
572
+ metadata = {"task_name": task_name, "doc_id": str(doc_id), "choice_idx": str(choice_idx)}
573
+ metadata.update(meta or {})
574
+ save_file({k: v.contiguous() for k, v in tensors.items()}, path, metadata=metadata)
575
+ self._count += 1
576
+ return path
577
+
578
+ @property
579
+ def files_written(self) -> int:
580
+ return self._count
581
+
582
+
583
+ class RunWriter:
584
+ """Owns every writer of one run directory and rolls all shards together.
585
+
586
+ Rolling every table at the same document boundary means shard N of `signals` and shard N of `steps` cover the same documents, which makes a partially written run easy to reason about.
587
+
588
+ Example:
589
+ >>> with RunWriter("/tmp/run") as w: # doctest: +SKIP
590
+ ... w.write_document({"steps": [...], "signals": [...]})
591
+ """
592
+
593
+ def __init__(self, run_dir: str | os.PathLike, shard_size: int = SHARD_SIZE) -> None:
594
+ self.run_dir = str(run_dir)
595
+ os.makedirs(self.run_dir, exist_ok=True)
596
+ acquire_writer_lock(self.run_dir)
597
+ self.shard_size = shard_size
598
+ # Read before opening any writer: which documents a restart must not record a second time.
599
+ self.already_recorded = existing_doc_keys(self.run_dir, "steps")
600
+ self.tables = {name: ShardWriter(self.run_dir, spec) for name, spec in TABLES.items()}
601
+ self.attention = TensorStore(self.run_dir, "attention")
602
+ self.raw_hidden = TensorStore(self.run_dir, "raw")
603
+ self._docs_since_flush = 0
604
+ self.documents_written = 0
605
+
606
+ def register_table(self, spec: TableSpec) -> None:
607
+ """실행별 custom schema를 등록하고 기존 schema와의 혼합을 거부한다."""
608
+ import base64
609
+ import re
610
+ if not re.fullmatch(r"custom/[A-Za-z][A-Za-z0-9_]*", spec.name):
611
+ raise ValueError("custom table must use custom/<name>")
612
+ schema = spec.arrow_schema()
613
+ path = os.path.join(self.run_dir, spec.name, "schema.json")
614
+ encoded = base64.b64encode(schema.serialize().to_pybytes()).decode("ascii")
615
+ if os.path.exists(path):
616
+ with open(path, encoding="utf-8") as stream:
617
+ previous = json.load(stream)
618
+ if previous["arrow_schema"] != encoded:
619
+ raise ValueError(f"custom schema changed: {spec.name}")
620
+ else:
621
+ os.makedirs(os.path.dirname(path), exist_ok=True)
622
+ with open(path, "w", encoding="utf-8") as stream:
623
+ json.dump({"arrow_schema": encoded, "metadata": json.loads(schema.metadata[b"eval_framework"])}, stream, indent=2)
624
+ if spec.name not in self.tables:
625
+ self.tables[spec.name] = ShardWriter(self.run_dir, spec)
626
+
627
+ def write_document(
628
+ self, rows_by_table: dict[str, list[dict[str, Any]]], documents: int = 1
629
+ ) -> None:
630
+ """Add one batch's worth of rows and roll the shards when due.
631
+
632
+ Args:
633
+ rows_by_table: table name -> rows.
634
+ Tables absent from the dict are simply not written for this document (e.g. `attn_norm` when --save-attention is off).
635
+ documents: how many documents these rows cover.
636
+ Greater than one when a batched forward pass finished several documents at once; it only affects when the next shard is rolled.
637
+ """
638
+ for table_name, rows in rows_by_table.items():
639
+ if table_name not in self.tables:
640
+ raise KeyError(f"unknown table {table_name!r}; declared tables: {list(self.tables)}")
641
+ self.tables[table_name].extend(rows)
642
+ self.documents_written += documents
643
+ self._docs_since_flush += documents
644
+ if self._docs_since_flush >= self.shard_size:
645
+ self.flush()
646
+
647
+ def flush(self) -> None:
648
+ """Close the current shard of every table and start the next one."""
649
+ for writer in self.tables.values():
650
+ writer.flush()
651
+ self._docs_since_flush = 0
652
+
653
+ def signal_files(self) -> list[str]:
654
+ """Relative paths of everything this run produced, for `results.json`."""
655
+ found: list[str] = []
656
+ custom_root = os.path.join(self.run_dir, "custom")
657
+ custom_tables = [f"custom/{name}" for name in os.listdir(custom_root)
658
+ if os.path.isdir(os.path.join(custom_root, name))] if os.path.isdir(custom_root) else []
659
+ for sub in dict.fromkeys(list(self.tables) + ["attention", "raw"] + custom_tables):
660
+ directory = os.path.join(self.run_dir, sub)
661
+ if not os.path.isdir(directory):
662
+ continue
663
+ if sub.startswith("custom/"):
664
+ from pathlib import Path
665
+ found.extend(str(p.relative_to(self.run_dir)) for p in Path(directory).iterdir() if p.is_file())
666
+ found.extend(str(p.relative_to(self.run_dir)) for p in (Path(directory) / "tensors").glob("*") if p.is_file() and p.suffix != ".tmp")
667
+ for marker in (Path(self.run_dir) / "custom_collection/commits").glob("*.json"):
668
+ data = _read_custom_commit(self.run_dir, marker)
669
+ found.extend(p for p in data["files"] if p.startswith(sub + "/"))
670
+ else:
671
+ for name in sorted(os.listdir(directory)):
672
+ found.append(f"{sub}/{name}")
673
+ return found
674
+
675
+ def close(self) -> None:
676
+ for writer in self.tables.values():
677
+ writer.close()
678
+
679
+ def __enter__(self) -> "RunWriter":
680
+ return self
681
+
682
+ def __exit__(self, *exc: object) -> None:
683
+ self.close()
684
+ release_writer_lock(self.run_dir)
685
+
686
+
687
+ # One lock per run directory held by this process; see `acquire_writer_lock`.
688
+ _WRITER_LOCKS: dict[str, Any] = {}
689
+
690
+
691
+ def acquire_writer_lock(run_dir: str | os.PathLike) -> None:
692
+ """Hold an exclusive lock on a run directory, refusing a second writing process.
693
+
694
+ Shard numbers are chosen from the files already on disk, so two processes writing one
695
+ directory would reuse them. Reopening a directory this process already holds is allowed;
696
+ `run` and `collect-research-data` release the lock when they return or fail.
697
+ """
698
+ import fcntl
699
+
700
+ path = os.path.realpath(str(run_dir))
701
+ if path in _WRITER_LOCKS:
702
+ return
703
+ handle = open(os.path.join(path, ".writer.lock"), "a")
704
+ try:
705
+ fcntl.flock(handle, fcntl.LOCK_EX | fcntl.LOCK_NB)
706
+ except BlockingIOError:
707
+ handle.close()
708
+ raise ValueError(
709
+ f"another process is writing to {run_dir}; a run directory takes one writer at a time"
710
+ ) from None
711
+ _WRITER_LOCKS[path] = handle
712
+
713
+
714
+ def release_writer_lock(run_dir: str | os.PathLike) -> None:
715
+ """Release this process's lock on a run directory, if it holds one."""
716
+ handle = _WRITER_LOCKS.pop(os.path.realpath(str(run_dir)), None)
717
+ if handle is not None:
718
+ handle.close()
719
+
720
+
721
+ def existing_doc_keys(run_dir: str | os.PathLike, table: str = "steps") -> set[tuple[str, int, int]]:
722
+ """Read back which (task_name, doc_id, choice_idx) triples are already stored.
723
+
724
+ Used to resume a run that died: whatever a completed shard contains does not need to be recomputed.
725
+
726
+ Example:
727
+ >>> existing_doc_keys("results/xnli_ko/qwen3-8b/2026-01-01-ab12cd") # doctest: +SKIP
728
+ {('xnli_ko', 0, 0), ('xnli_ko', 0, 1), ('xnli_ko', 1, 0)}
729
+ """
730
+ directory = os.path.join(str(run_dir), table)
731
+ if not os.path.isdir(directory):
732
+ return set()
733
+ keys: set[tuple[str, int, int]] = set()
734
+ for name in sorted(os.listdir(directory)):
735
+ if not name.endswith(".parquet"):
736
+ continue
737
+ path = os.path.join(directory, name)
738
+ try:
739
+ shard = pq.read_table(path, columns=["task_name", "doc_id", "choice_idx"])
740
+ except (pa.ArrowException, OSError) as error:
741
+ raise ValueError(f"cannot read recorded shard {path}: {error}") from error
742
+ for task_name, doc_id, choice_idx in zip(
743
+ shard.column("task_name").to_pylist(),
744
+ shard.column("doc_id").to_pylist(),
745
+ shard.column("choice_idx").to_pylist(),
746
+ ):
747
+ keys.add((task_name, int(doc_id), int(choice_idx)))
748
+ return keys
749
+
750
+
751
+ def check_recorded_tables_agree(run_dir: str | os.PathLike, tables: Iterable[str]) -> None:
752
+ """Refuse to resume when per-document tables on disk cover different documents.
753
+
754
+ Resume skips every (task_name, doc_id, choice_idx) already in `steps`. A missing shard,
755
+ or a flush that stopped between tables, would otherwise resume as silently missing
756
+ signal rows, or as a second copy of them.
757
+ """
758
+ steps = existing_doc_keys(run_dir, "steps")
759
+ for table in tables:
760
+ keys = existing_doc_keys(run_dir, table)
761
+ if keys != steps:
762
+ raise ValueError(
763
+ f"{run_dir}: {table} covers {len(keys)} (task, doc, choice) keys but steps covers "
764
+ f"{len(steps)} ({len(steps - keys)} missing from {table}, {len(keys - steps)} absent "
765
+ "from steps). A shard is missing or a write stopped between tables, so resuming "
766
+ "would drop or duplicate rows. Restore the shards or use a new output directory."
767
+ )
768
+
769
+
770
+ # --------------------------------------------------------------------------
771
+ # Manifest and results.json
772
+ # --------------------------------------------------------------------------
773
+
774
+ RESULTS_FILENAME = "results.json"
775
+ SAMPLES_FILENAME = "samples.jsonl"
776
+
777
+ # : Keys the manifest must carry.
778
+ # Two runs whose manifests : differ in any of these are not comparable, so a missing key would let : `report.py` group runs it should have kept apart.
779
+ REQUIRED_MANIFEST_KEYS = (
780
+ "schema_version",
781
+ "tool_version",
782
+ "lm_eval_version",
783
+ "model_id",
784
+ "tokenizer_id",
785
+ "tasks",
786
+ "num_fewshot",
787
+ "limit",
788
+ "doc_id_set_hash",
789
+ "reducers",
790
+ "n_blocks",
791
+ "layer_index_convention",
792
+ "attn_index_convention",
793
+ "fixed_settings",
794
+ "seed",
795
+ "started_at",
796
+ "completed",
797
+ )
798
+
799
+
800
+ def write_results(
801
+ run_dir: str | os.PathLike,
802
+ manifest: dict[str, Any],
803
+ scores: dict[str, Any],
804
+ signal_files: list[str] | None = None,
805
+ ) -> str:
806
+ """Write `results.json`: lm-eval's scores plus our manifest plus a file list.
807
+
808
+ The manifest is what `report.py` reads; it never parses the directory path, because `--output` lets the user put a run anywhere.
809
+
810
+ Example:
811
+ >>> write_results("/tmp/run", manifest, {"xnli_ko": {"acc": 0.71}}) # doctest: +SKIP
812
+ '/tmp/run/results.json'
813
+ """
814
+ missing = [k for k in REQUIRED_MANIFEST_KEYS if k not in manifest]
815
+ if missing:
816
+ raise ValueError(f"manifest is missing required keys: {missing}")
817
+ payload = {
818
+ "schema_version": SCHEMA_VERSION,
819
+ "manifest": manifest,
820
+ "results": scores,
821
+ "signal_files": signal_files or [],
822
+ }
823
+ os.makedirs(str(run_dir), exist_ok=True)
824
+ path = os.path.join(str(run_dir), RESULTS_FILENAME)
825
+ _atomic_json(path, payload, ensure_ascii=False, default=str, allow_nan=True)
826
+ return path
827
+
828
+
829
+ def read_results(run_dir: str | os.PathLike) -> dict[str, Any]:
830
+ """Load `results.json`.
831
+ Raises FileNotFoundError if this is not a run directory."""
832
+ with open(os.path.join(str(run_dir), RESULTS_FILENAME), encoding="utf-8") as fh:
833
+ return json.load(fh)
834
+
835
+
836
+ def read_manifest(run_dir: str | os.PathLike) -> dict[str, Any]:
837
+ """Load only the manifest part of `results.json`.
838
+
839
+ Example:
840
+ >>> read_manifest("results/xnli_ko/qwen3-8b/2026-01-01-ab12cd")["n_blocks"] # doctest: +SKIP
841
+ 36
842
+ """
843
+ return read_results(run_dir)["manifest"]
844
+
845
+
846
+ def mark_complete(run_dir: str | os.PathLike, extra: dict[str, Any] | None = None) -> None:
847
+ """Stamp a run as finished.
848
+
849
+ `report.py` excludes runs without this marker: a run that died mid-way has a truncated document set, and averaging it against a complete run is a silent mistake.
850
+ """
851
+ payload = read_results(run_dir)
852
+ payload["manifest"]["completed"] = True
853
+ payload["manifest"]["finished_at"] = datetime.now(timezone.utc).isoformat()
854
+ payload["manifest"].update(extra or {})
855
+ _atomic_json(os.path.join(str(run_dir), RESULTS_FILENAME), payload,
856
+ ensure_ascii=False, default=str, allow_nan=True)
857
+
858
+
859
+ # --------------------------------------------------------------------------
860
+ # samples.jsonl (lm-eval --log_samples, kept verbatim inside the run dir)
861
+ # --------------------------------------------------------------------------
862
+
863
+
864
+ def write_samples(run_dir: str | os.PathLike, samples_by_task: dict[str, list[dict]]) -> str:
865
+ """Persist lm-eval's per-sample log inside the run directory.
866
+
867
+ We write it ourselves from `simple_evaluate(..., log_samples=True)` rather than letting lm-eval's EvaluationTracker do it, for two reasons: the tracker writes outside our run directory, and it rewrites `arguments` into a flattened dict, which `collect-research-data` needs in its original (context, continuation) form.
868
+
869
+ Example of one written line:
870
+ {"task_name": "xnli_ko", "doc_id": 0, "target": 1, "arguments": [["...ctx...", " Yes"], ["...ctx...", " No"]], "filtered_resps": [[-3.1, false], [-2.4, false]], "acc": 1.0, "prompt_hash": "d41d8c..."}
871
+ """
872
+ path = os.path.join(str(run_dir), SAMPLES_FILENAME)
873
+ os.makedirs(str(run_dir), exist_ok=True)
874
+ # Judge grading checkpoints rewrite this log. Keep the previous complete
875
+ # generation log if serialization or writing the replacement fails.
876
+ temporary = path + '.tmp'
877
+ with open(temporary, "w", encoding="utf-8") as fh:
878
+ for task_name, samples in samples_by_task.items():
879
+ for sample in samples:
880
+ record = dict(sample)
881
+ record["task_name"] = task_name
882
+ # `arguments` holds tuples; json turns them into lists anyway, but we normalise here so collection always sees the same shape.
883
+ record["arguments"] = [list(arg) for arg in record.get("arguments", [])]
884
+ fh.write(json.dumps(record, ensure_ascii=False, default=str) + "\n")
885
+ os.replace(temporary, path)
886
+ return path
887
+
888
+
889
+ def read_samples(run_dir: str | os.PathLike) -> list[dict[str, Any]]:
890
+ """Read `samples.jsonl` back as a list of records.
891
+
892
+ Example:
893
+ >>> samples = read_samples("results/xnli_ko/qwen3-8b/2026-01-01-ab12cd") # doctest: +SKIP
894
+ >>> samples[0]["task_name"], samples[0]["doc_id"]
895
+ ('xnli_ko', 0)
896
+ """
897
+ path = os.path.join(str(run_dir), SAMPLES_FILENAME)
898
+ with open(path, encoding="utf-8") as fh:
899
+ return [json.loads(line) for line in fh if line.strip()]
900
+
901
+
902
+ # --------------------------------------------------------------------------
903
+ # Joining lm-eval's grading result onto the `docs` table
904
+ # --------------------------------------------------------------------------
905
+
906
+ # : Metrics we treat as a per-document right/wrong verdict, in preference order. : We never re-derive correctness ourselves; we only read what lm-eval scored.
907
+ _CORRECTNESS_METRICS = ("acc", "exact_match", "acc_norm", "em", "f1")
908
+
909
+
910
+ def _extract_is_correct(sample: dict[str, Any]) -> bool | None:
911
+ """Pull a boolean verdict out of an lm-eval sample record.
912
+
913
+ lm-eval stores the metric values it computed directly on the record (`sample["acc"] = 1.0`).
914
+ We take the first metric we recognise and treat 1.0 as correct.
915
+ Returns None when the task reports no such metric (e.g. a purely generative task scored by ROUGE), in which case `is_correct` stays null rather than being guessed.
916
+
917
+ Example:
918
+ >>> _extract_is_correct({"metrics": ["acc"], "acc": 1.0})
919
+ True
920
+ >>> _extract_is_correct({"metrics": ["rouge1"], "rouge1": 0.42}) is None
921
+ True
922
+ """
923
+ if "_eval_framework" in sample:
924
+ return sample["_eval_framework"]["is_correct"]
925
+ for metric in _CORRECTNESS_METRICS:
926
+ if metric in sample:
927
+ try:
928
+ return float(sample[metric]) == 1.0
929
+ except (TypeError, ValueError):
930
+ return None
931
+ return None
932
+
933
+
934
+ def _gold_choice_index(
935
+ sample: dict[str, Any], n_choices: int, choices: Sequence[str] | None = None
936
+ ) -> int | None:
937
+ """Work out which choice was the gold one.
938
+
939
+ Two forms, because lm-eval accepts two and resolves them the same way itself (`api/task.py`: `gold = choices.index(gold) if isinstance(gold, str)`).
940
+
941
+ * an **index**, which is what `doc_to_target` yields for a task whose target field is a number. It is stringified on the way to disk, hence the int() attempt.
942
+ * a **label**, for a task whose `doc_to_target` names the answer and whose `doc_to_choice` lists the labels - `doc_to_target: answer` over `["A", "B", "C", "D"]`, which is how Global-MMLU and most letter-choice benchmarks are written. The target is then `"A"`, and reading it as an index gives nothing.
943
+
944
+ Matching a label needs the choices, which are the continuations of the logged request, compared with surrounding whitespace removed: lm-eval joins a choice to its context with a delimiter (a space by default), so the continuation on disk is `" A"` where the target is `"A"`.
945
+
946
+ Anything still unmatched - a free-form target, an out-of-range index - yields None, and `is_target_choice` is then false for every choice rather than being invented.
947
+
948
+ Example:
949
+ >>> _gold_choice_index({"target": "1"}, n_choices=3)
950
+ 1
951
+ >>> _gold_choice_index({"target": "C"}, 4, choices=[" A", " B", " C", " D"])
952
+ 2
953
+ >>> _gold_choice_index({"target": "Paris"}, n_choices=3) is None
954
+ True
955
+ """
956
+ target = sample.get("target")
957
+ try:
958
+ index = int(target)
959
+ except (TypeError, ValueError):
960
+ pass
961
+ else:
962
+ return index if 0 <= index < n_choices else None
963
+ if choices and isinstance(target, str):
964
+ wanted = target.strip()
965
+ for position, choice in enumerate(choices):
966
+ if str(choice).strip() == wanted:
967
+ return position if position < n_choices else None
968
+ return None
969
+
970
+
971
+ def available_filters(samples: list[dict[str, Any]]) -> dict[str, list[str]]:
972
+ """Which lm-eval filters appear per task, in order of first appearance.
973
+
974
+ Example:
975
+ >>> available_filters([{"task_name": "gsm8k", "filter": "strict-match"},
976
+ ... {"task_name": "gsm8k", "filter": "flexible-extract"},
977
+ ... {"task_name": "gsm8k", "filter": "strict-match"}])
978
+ {'gsm8k': ['strict-match', 'flexible-extract']}
979
+ """
980
+ found: dict[str, list[str]] = {}
981
+ for sample in samples:
982
+ name = sample.get("filter")
983
+ if name is None:
984
+ continue
985
+ seen = found.setdefault(sample["task_name"], [])
986
+ if name not in seen:
987
+ seen.append(name)
988
+ return found
989
+
990
+
991
+ def primary_filters(samples: list[dict[str, Any]]) -> dict[str, str]:
992
+ """The one filter per task whose verdict goes into `docs`.
993
+
994
+ lm-eval emits one sample record per filter, so a task with two filters would otherwise give `docs` two rows per document - and since `load_signals` joins on (task_name, doc_id, choice_idx), that would silently double every signal row.
995
+ The first filter lm-eval reports is used, which is its own declaration order and therefore stable across runs of the same task.
996
+
997
+ Example:
998
+ >>> primary_filters([{"task_name": "gsm8k", "filter": "strict-match"},
999
+ ... {"task_name": "gsm8k", "filter": "flexible-extract"}])
1000
+ {'gsm8k': 'strict-match'}
1001
+ """
1002
+ primary = {task: names[0] for task, names in available_filters(samples).items()}
1003
+ for sample in samples:
1004
+ if "_eval_framework" in sample:
1005
+ primary[sample["task_name"]] = sample["_eval_framework"]["primary_filter"]
1006
+ return primary
1007
+
1008
+
1009
+ def build_docs_rows(samples: list[dict[str, Any]]) -> list[dict[str, Any]]:
1010
+ """Turn lm-eval sample records into `docs` rows, one per (doc, choice).
1011
+
1012
+ Grading stays lm-eval's job here: and all this does is reshape its output from per-document to per-(document, choice) so it lines up with the signal tables.
1013
+
1014
+ Records for a non-primary filter are skipped, so the result stays unique on (task_name, doc_id, choice_idx).
1015
+
1016
+ Example:
1017
+ >>> rows = build_docs_rows([{
1018
+ ... "task_name": "xnli_ko", "doc_id": 7, "target": "1",
1019
+ ... "arguments": [["ctx", " Yes"], ["ctx", " No"]],
1020
+ ... "filtered_resps": [[-3.1, False], [-2.4, False]],
1021
+ ... "metrics": ["acc"], "acc": 1.0}])
1022
+ >>> [(r["choice_idx"], r["is_target_choice"], r["predicted"]) for r in rows]
1023
+ [(0, False, ' No'), (1, True, ' No')]
1024
+ """
1025
+ primary = primary_filters(samples)
1026
+ rows: list[dict[str, Any]] = []
1027
+ for sample in samples:
1028
+ task_name = sample["task_name"]
1029
+ sample_filter = sample.get("filter")
1030
+ if sample_filter is not None and primary.get(task_name) != sample_filter:
1031
+ continue
1032
+ doc_id = int(sample["doc_id"])
1033
+ is_correct = _extract_is_correct(sample)
1034
+ responses = sample.get("filtered_resps") or []
1035
+ arguments = sample.get("arguments") or []
1036
+
1037
+ # A loglikelihood response is [logprob, is_greedy]; a generate response is a plain string.
1038
+ # The shape tells us which kind of task this is.
1039
+ is_loglikelihood = bool(responses) and isinstance(responses[0], (list, tuple))
1040
+
1041
+ if is_loglikelihood:
1042
+ logprobs = [float(resp[0]) for resp in responses]
1043
+ continuations = [
1044
+ str(argument[1]) for argument in arguments if len(argument) > 1]
1045
+ gold = _gold_choice_index(sample, len(logprobs), continuations)
1046
+ best = max(range(len(logprobs)), key=logprobs.__getitem__)
1047
+ # Report the winning continuation string when we have it, so the column is readable without cross-referencing the dataset.
1048
+ predicted = (
1049
+ str(arguments[best][1]) if best < len(arguments) and len(arguments[best]) > 1
1050
+ else str(best)
1051
+ )
1052
+ target_text = (
1053
+ str(arguments[gold][1]) if gold is not None and gold < len(arguments)
1054
+ and len(arguments[gold]) > 1 else str(sample.get("target"))
1055
+ )
1056
+ for choice_idx, logprob in enumerate(logprobs):
1057
+ rows.append(
1058
+ {
1059
+ "task_name": task_name,
1060
+ "doc_id": doc_id,
1061
+ "choice_idx": choice_idx,
1062
+ "is_correct": is_correct,
1063
+ "is_target_choice": gold is not None and choice_idx == gold,
1064
+ "predicted": predicted,
1065
+ "target": target_text,
1066
+ "choice_logprob": logprob,
1067
+ "filter": sample_filter,
1068
+ }
1069
+ )
1070
+ else:
1071
+ # A generate task has one result row per document, with choice_idx=0.
1072
+ # Generation may use many forward passes; this table stores the final
1073
+ # response and has no per-choice log probability.
1074
+ rows.append(
1075
+ {
1076
+ "task_name": task_name,
1077
+ "doc_id": doc_id,
1078
+ "choice_idx": 0,
1079
+ "is_correct": is_correct,
1080
+ "is_target_choice": True,
1081
+ "predicted": str(responses[0]) if responses else "",
1082
+ "target": str(sample.get("target")),
1083
+ "choice_logprob": None,
1084
+ "filter": sample_filter,
1085
+ }
1086
+ )
1087
+ return rows
1088
+
1089
+
1090
+ def write_docs_table(run_dir: str | os.PathLike, samples: list[dict[str, Any]]) -> str | None:
1091
+ """Write the whole `docs` table in one shot from `samples.jsonl`.
1092
+
1093
+ Done at the end of a run, not afterwards by a separate script: if the join lived outside the run, the run directory alone would not be reproducible.
1094
+
1095
+ Tasks with more than one filter are reduced to the primary one and say so on stderr, because which filter `is_correct` came from changes what a conditional analysis means.
1096
+
1097
+ Unlike the signal tables this one is rebuilt in full from `samples.jsonl` every time, so any earlier shards are removed first.
1098
+ Adding to them would duplicate the key that `load_signals` joins on.
1099
+ """
1100
+ import sys
1101
+
1102
+ directory = os.path.join(str(run_dir), "docs")
1103
+ if os.path.isdir(directory):
1104
+ for name in os.listdir(directory):
1105
+ if name.endswith(".parquet"):
1106
+ os.remove(os.path.join(directory, name))
1107
+
1108
+ for task, names in available_filters(samples).items():
1109
+ if len(names) > 1:
1110
+ print(
1111
+ f"note: task {task} defines filters {names}; `docs` carries "
1112
+ f"{names[0]!r}. The others remain in {SAMPLES_FILENAME}.",
1113
+ file=sys.stderr,
1114
+ )
1115
+ writer = ShardWriter(run_dir, TABLES["docs"])
1116
+ writer.extend(build_docs_rows(samples))
1117
+ return writer.flush()
1118
+
1119
+
1120
+ # --------------------------------------------------------------------------
1121
+ # Reading
1122
+ # --------------------------------------------------------------------------
1123
+
1124
+
1125
+ def read_table(run_dir: str | os.PathLike, table: str):
1126
+ """Read every shard of one table as a pandas DataFrame (empty if absent).
1127
+
1128
+ Example:
1129
+ >>> read_table(run_dir, "signals").columns.tolist() # doctest: +SKIP
1130
+ ['task_name', 'doc_id', 'choice_idx', 'step', 'layer', 'lens_token_id', ...]
1131
+ """
1132
+ import pandas as pd
1133
+
1134
+ if table.startswith("custom/"):
1135
+ import base64
1136
+ import re
1137
+ if not re.fullmatch(r"custom/[A-Za-z][A-Za-z0-9_]*", table):
1138
+ raise ValueError("invalid custom table name")
1139
+ directory = os.path.join(str(run_dir), table)
1140
+ with open(os.path.join(directory, "schema.json"), encoding="utf-8") as stream:
1141
+ stored = json.load(stream)
1142
+ schema = pa.ipc.read_schema(pa.BufferReader(base64.b64decode(stored["arrow_schema"])))
1143
+ shards = sorted(os.path.join(directory, n) for n in os.listdir(directory) if n.endswith(".parquet"))
1144
+ from pathlib import Path
1145
+ for marker in (Path(run_dir) / "custom_collection" / "commits").glob("*.json"):
1146
+ data = _read_custom_commit(run_dir, marker)
1147
+ if table in data["tables"]:
1148
+ shards.append(str(Path(run_dir) / data["tables"][table]))
1149
+ return pq.read_table(sorted(shards), schema=schema).to_pandas() if shards else pd.DataFrame(columns=schema.names)
1150
+ spec = TABLES[table]
1151
+ directory = os.path.join(str(run_dir), table)
1152
+ if not os.path.isdir(directory):
1153
+ return pd.DataFrame(columns=spec.column_names)
1154
+ shards = sorted(
1155
+ os.path.join(directory, name)
1156
+ for name in os.listdir(directory)
1157
+ if name.endswith(".parquet")
1158
+ )
1159
+ if not shards:
1160
+ return pd.DataFrame(columns=spec.column_names)
1161
+ return pq.read_table(shards).to_pandas()
1162
+
1163
+
1164
+ def load_signals(path: str | os.PathLike, table: str = "signals", join_docs: bool = True):
1165
+ """Load one signal table of one run, with the grading columns joined in.
1166
+
1167
+ Grading is stored once in `docs` and joined at read time rather than being copied into every signal row: `signals` has (layers x steps x choices) rows per document, so duplicating the answer strings there would bloat the files for no gain.
1168
+
1169
+ Args:
1170
+ path: A run directory (the one containing `results.json`).
1171
+ table: Which signal table to load - "signals", "similarity", "steps" or "attn_norm".
1172
+ join_docs: Attach `is_correct` / `is_target_choice` / `predicted` / `target` / `choice_logprob` from the `docs` table.
1173
+
1174
+ Returns:
1175
+ A pandas DataFrame.
1176
+
1177
+ Example:
1178
+ >>> df = load_signals("results/xnli_ko/qwen3-8b/2026-01-01-ab12cd") # doctest: +SKIP
1179
+ >>> df.query("is_correct and layer == 20").lens_prob.mean()
1180
+ 0.42
1181
+ """
1182
+ frame = read_table(path, table)
1183
+ if join_docs and table != "docs" and len(frame):
1184
+ docs = read_table(path, "docs")
1185
+ if len(docs):
1186
+ # A duplicated key here would multiply every signal row instead of annotating it, and the result looks like ordinary data.
1187
+ # Refuse.
1188
+ key = ["task_name", "doc_id", "choice_idx"]
1189
+ duplicated = docs.duplicated(subset=key).sum()
1190
+ if duplicated:
1191
+ raise ValueError(
1192
+ f"`docs` has {duplicated} rows duplicating its key {key}; joining "
1193
+ "would silently multiply every signal row. This usually means "
1194
+ "several lm-eval filters were written to `docs` instead of one."
1195
+ )
1196
+ frame = frame.merge(docs, on=key, how="left")
1197
+ return frame
1198
+
1199
+
1200
+ def signal_columns_report(table: str = "signals") -> str:
1201
+ """The per-column documentation for one table, as printed next to the data.
1202
+
1203
+ Example:
1204
+ >>> print(signal_columns_report("similarity")) # doctest: +ELLIPSIS
1205
+ similarity: Pairwise cosine similarity...
1206
+ layer_i int32 ...
1207
+ """
1208
+ spec = TABLES[table]
1209
+ lines = [f"{spec.name}: {spec.description}"]
1210
+ for col in spec.columns:
1211
+ vocab = "any-vocab" if col.comparable_across_vocab else "same-vocab-only"
1212
+ lines.append(f" {col.name:<18} {str(col.dtype):<8} {vocab:<15} {col.description}")
1213
+ return "\n".join(lines)
1214
+
1215
+
1216
+ def read_sample_metrics(run_dir: str | os.PathLike, *, metric: str, filter_name: str):
1217
+ """Read exactly one metric/filter per document, without duplicating signal joins.
1218
+
1219
+ Returns a DataFrame keyed by task_name/doc_id with sample_id and value. Values
1220
+ remain objects: corpus metrics may store tuples rather than scalar scores.
1221
+ Join with signals using ``validate="many_to_one"`` on task_name/doc_id.
1222
+ """
1223
+ import pandas as pd
1224
+
1225
+ rows, seen = [], set()
1226
+ for sample in read_samples(run_dir):
1227
+ if sample.get("filter", "none") != filter_name or metric not in sample.get("metrics", []):
1228
+ continue
1229
+ key = (sample["task_name"], int(sample["doc_id"]))
1230
+ if key in seen:
1231
+ raise ValueError(f"duplicate metric/filter sample: {key}")
1232
+ seen.add(key)
1233
+ rows.append({"task_name": key[0], "doc_id": key[1],
1234
+ "sample_id": sample.get("_eval_framework", {}).get("sample_id"),
1235
+ "value": sample[metric]})
1236
+ return pd.DataFrame(rows, columns=["task_name", "doc_id", "sample_id", "value"])
1237
+
1238
+
1239
+ def _custom_checkpoint(phase):
1240
+ """Fault-injection seam for subprocess durability tests; production is a no-op."""
1241
+
1242
+
1243
+ # Custom tensors deliberately use unique artifacts, not TensorStore.path_for(): that
1244
+ # legacy path identifies only a document/choice and overwrites repeated calls.
1245
+ def _sync_directory(path):
1246
+ """Persist an atomic rename's directory entry on supported local filesystems."""
1247
+ descriptor = os.open(str(path), os.O_RDONLY)
1248
+ try:
1249
+ os.fsync(descriptor)
1250
+ finally:
1251
+ os.close(descriptor)
1252
+
1253
+
1254
+ def _atomic_json(path, value, **json_options):
1255
+ """Publish metadata only after its complete bytes are durable on the local filesystem."""
1256
+ from pathlib import Path
1257
+ import uuid
1258
+ path = Path(path)
1259
+ path.parent.mkdir(parents=True, exist_ok=True)
1260
+ temporary = path.with_name(path.name + "." + uuid.uuid4().hex + ".tmp")
1261
+ with temporary.open("w", encoding="utf-8") as stream:
1262
+ options = {"sort_keys": True, "indent": 2, "allow_nan": False, **json_options}
1263
+ json.dump(value, stream, **options)
1264
+ stream.flush()
1265
+ os.fsync(stream.fileno())
1266
+ os.replace(temporary, path)
1267
+ _sync_directory(path.parent)
1268
+
1269
+
1270
+ def save_custom_tensors(run_dir, name, tensors, rows, output_features, attempt=None, max_bytes=67108864, input_features="input_features"):
1271
+ """Write an immutable safetensors file, then publish its independent JSON index.
1272
+
1273
+ Limits include the serialized header. A crash between the two publications leaves
1274
+ an orphan artifact, detected by audit_custom_tensors; it is never silently indexed.
1275
+ """
1276
+ from pathlib import Path
1277
+ import uuid
1278
+ import hashlib
1279
+ from safetensors.torch import save_file
1280
+ import re
1281
+ if not re.fullmatch(r"[A-Za-z][A-Za-z0-9_]*", name) or (attempt is not None and not re.fullmatch(r"[a-f0-9]{32}", attempt)):
1282
+ raise ValueError("invalid custom tensor namespace")
1283
+ root = Path(run_dir)
1284
+ namespace = root / "custom" / name
1285
+ if attempt is not None:
1286
+ namespace = namespace / "attempts" / attempt
1287
+ directory = namespace / "tensors"
1288
+ directory.mkdir(parents=True, exist_ok=True)
1289
+ identifier = uuid.uuid4().hex
1290
+ path = directory / (identifier + ".safetensors")
1291
+ temporary = path.with_suffix(".tmp")
1292
+ _custom_checkpoint("tensor_before")
1293
+ save_file({k: t.detach().cpu().contiguous() for k, t in tensors.items()}, str(temporary))
1294
+ size = temporary.stat().st_size
1295
+ if size > max_bytes:
1296
+ temporary.unlink()
1297
+ raise ValueError("custom max_tensor_bytes exceeded (serialized file)")
1298
+ with temporary.open("rb") as stream:
1299
+ os.fsync(stream.fileno())
1300
+ os.replace(temporary, path)
1301
+ _sync_directory(path.parent)
1302
+ _custom_checkpoint("tensor_after")
1303
+ index = {"schema_version": 1, "tensor_id": identifier, "path": str(path.relative_to(root)),
1304
+ "bytes": size, "sha256": hashlib.sha256(path.read_bytes()).hexdigest(), "rows": rows,
1305
+ "tensors": {k: {"dtype": str(t.dtype), "shape": list(t.shape),
1306
+ "axes": ["selected_input_position", output_features if k == "output" else input_features]}
1307
+ for k, t in tensors.items()}}
1308
+ index_path = path.with_suffix(".json")
1309
+ _custom_checkpoint("index_before")
1310
+ _atomic_json(index_path, index)
1311
+ _custom_checkpoint("index_after")
1312
+ return str(index_path.relative_to(root)), size
1313
+
1314
+
1315
+ def _custom_index_paths(run_dir, name):
1316
+ from pathlib import Path
1317
+ root = Path(run_dir)
1318
+ paths = list((root / "custom" / name / "tensors").glob("*.json"))
1319
+ for marker in (root / "custom_collection" / "commits").glob("*.json"):
1320
+ data = _read_custom_commit(root, marker)
1321
+ paths.extend(root / p for p in data["indexes"] if p.startswith(f"custom/{name}/"))
1322
+ return sorted(paths)
1323
+
1324
+
1325
+ def read_custom_tensors(run_dir, name):
1326
+ """Return (index, tensor dict) pairs; no user factory or arbitrary pickle is imported.
1327
+
1328
+ By default collection attempts are visible only through their committed marker.
1329
+ Content hashes and declared dtype/shape are checked before returning values.
1330
+ """
1331
+ from pathlib import Path
1332
+ import hashlib
1333
+ from safetensors.torch import load_file
1334
+ root = Path(run_dir)
1335
+ result = []
1336
+ for index_path in _custom_index_paths(root, name):
1337
+ data = json.loads(index_path.read_text())
1338
+ path = root / data["path"]
1339
+ if not path.resolve().is_relative_to(root.resolve()):
1340
+ raise ValueError("custom tensor path escapes run directory")
1341
+ if not path.is_file() or hashlib.sha256(path.read_bytes()).hexdigest() != data["sha256"]:
1342
+ raise ValueError("custom tensor artifact missing or corrupt")
1343
+ tensors = load_file(str(path))
1344
+ if set(tensors) != set(data["tensors"]):
1345
+ raise ValueError("custom tensor index keys mismatch")
1346
+ for key, tensor in tensors.items():
1347
+ declared = data["tensors"][key]
1348
+ if str(tensor.dtype) != declared["dtype"] or list(tensor.shape) != declared["shape"]:
1349
+ raise ValueError("custom tensor dtype/shape mismatch")
1350
+ result.append((data, tensors))
1351
+ return result
1352
+
1353
+
1354
+ def audit_custom_tensors(run_dir, name):
1355
+ """Report orphan artifacts and missing artifacts, including uncommitted attempts."""
1356
+ from pathlib import Path
1357
+ root = Path(run_dir)
1358
+ base = root / "custom" / name
1359
+ artifacts = set(base.rglob("*.safetensors"))
1360
+ indexes = list(base.rglob("tensors/*.json"))
1361
+ indexed = {root / json.loads(p.read_text())["path"] for p in indexes}
1362
+ return {"orphan_artifacts": sorted(str(p.relative_to(root)) for p in artifacts - indexed),
1363
+ "missing_artifacts": sorted(str(p.relative_to(root)) for p in indexed - artifacts),
1364
+ "temporary_files": sorted(str(p.relative_to(root)) for p in base.rglob("*.tmp"))}
1365
+
1366
+
1367
+ def _read_custom_commit(root, marker):
1368
+ """완료 marker의 구조·관측 범위·파일을 검증한 뒤 반환한다.
1369
+
1370
+ marker가 있어도 참조 파일이 손상되면 완료된 수집으로 읽지 않는다.
1371
+ 검사 순서를 유지해 여러 문제가 있을 때도 같은 오류를 먼저 보고한다.
1372
+ """
1373
+ from pathlib import Path
1374
+ root = Path(root)
1375
+ data = json.loads(Path(marker).read_text())
1376
+ required = {'version', 'key', 'attempt', 'tables', 'indexes', 'files', 'hooks'}
1377
+ if set(data) != required or data['version'] != 1:
1378
+ raise ValueError('incomplete custom collection marker')
1379
+ contract_path = root / 'custom_collection' / 'contract.json'
1380
+ if not contract_path.is_file():
1381
+ raise ValueError('incomplete custom collection manifest: missing contract')
1382
+ contract = json.loads(contract_path.read_text())
1383
+ if set(contract) != {'version', 'specs', 'resolved', 'config_identity', 'samples_sha256'} or contract['version'] != 1:
1384
+ raise ValueError('incomplete custom collection manifest')
1385
+ _validate_custom_coverage(data, contract)
1386
+ _validate_custom_files(root, data)
1387
+ return data
1388
+
1389
+
1390
+ def _validate_custom_coverage(data: dict[str, Any], contract: dict[str, Any]) -> None:
1391
+ """계약의 모든 hook·모듈이 marker에 있는지 확인한다. 호출 0회도 유효한 기록이다.
1392
+
1393
+ 호출 횟수는 음이 아닌 int만 허용한다. bool을 횟수로 받아들이지 않도록
1394
+ isinstance 대신 정확한 타입을 검사한다.
1395
+ """
1396
+ expected = contract['resolved']
1397
+ if set(data['hooks']) != set(expected) or set(data['tables']) != {f'custom/{n}' for n in expected}:
1398
+ raise ValueError('incomplete custom collection hook coverage')
1399
+ for name, paths in expected.items():
1400
+ calls = data['hooks'][name]
1401
+ if set(calls) != set(paths) or any(type(count) is not int or count < 0 for count in calls.values()):
1402
+ raise ValueError('incomplete custom collection call coverage')
1403
+
1404
+
1405
+ def _validate_custom_files(root, data: dict[str, Any]) -> None:
1406
+ """scalar·index·tensor 참조가 정확히 같은 파일 집합을 가리키는지 검사한다.
1407
+
1408
+ 먼저 marker가 나열한 모든 파일의 경로와 hash를 확인한다. 그다음 JSON index를
1409
+ 읽어 tensor 참조까지 대조한다. 손상된 index를 먼저 해석하지 않는 순서다.
1410
+ """
1411
+ import hashlib
1412
+
1413
+ references = list(data['tables'].values()) + data['indexes']
1414
+ if any(path not in data['files'] for path in references):
1415
+ raise ValueError('incomplete custom collection references')
1416
+ for relative, digest in data['files'].items():
1417
+ path = root / relative
1418
+ if not path.resolve().is_relative_to(root.resolve()):
1419
+ raise ValueError('custom collection path escapes run directory')
1420
+ if not path.is_file() or hashlib.sha256(path.read_bytes()).hexdigest() != digest:
1421
+ raise ValueError(f'committed custom collection file missing or corrupt: {relative}')
1422
+ indexed_artifacts = []
1423
+ for relative in data['indexes']:
1424
+ index = json.loads((root / relative).read_text())
1425
+ indexed_artifacts.append(index['path'])
1426
+ if data['files'].get(index['path']) != index['sha256']:
1427
+ raise ValueError('incomplete custom collection tensor references')
1428
+ if set(data['files']) != set(references + indexed_artifacts):
1429
+ raise ValueError('custom collection file set mismatch')
1430
+
1431
+
1432
+ class CustomCollectionStore:
1433
+ """One committed unit is (task, document, choice, collection pass).
1434
+
1435
+ Each attempt has unique scalar/tensor/index paths. The only publication point is
1436
+ the final marker, after all outputs and hashes exist. Interrupted attempts remain
1437
+ available for diagnosis and never become default reader output. A process lock
1438
+ prevents two collectors from racing to publish the same unit.
1439
+ """
1440
+ def __init__(self, run_dir, contract):
1441
+ from pathlib import Path
1442
+ import fcntl
1443
+ self.root = Path(run_dir)
1444
+ self.directory = self.root / 'custom_collection'
1445
+ self.directory.mkdir(exist_ok=True)
1446
+ self.lock = (self.directory / 'writer.lock').open('a')
1447
+ try:
1448
+ fcntl.flock(self.lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
1449
+ self._ensure_contract(contract)
1450
+ # 파일 존재만으로 완료를 추정하지 않는다. 각 marker의 참조를 검증한
1451
+ # 단위만 completed에 넣어 다음 수집에서 건너뛸 수 있게 한다.
1452
+ self.completed = {}
1453
+ for marker in (self.directory / 'commits').glob('*.json'):
1454
+ data = _read_custom_commit(self.root, marker)
1455
+ key = tuple(data['key'])
1456
+ if key in self.completed or marker.stem != self.key_id(key):
1457
+ raise ValueError('duplicate or misidentified custom collection commit')
1458
+ if set(data['hooks']) != set(contract['resolved']):
1459
+ raise ValueError('incomplete custom collection hook coverage')
1460
+ self.completed[key] = data
1461
+ except BaseException:
1462
+ self.close()
1463
+ raise
1464
+
1465
+ def _ensure_contract(self, contract: dict[str, Any]) -> None:
1466
+ """writer lock을 잡은 상태에서 기존 계약과 비교하거나 최초 계약을 저장한다.
1467
+
1468
+ JSON 왕복으로 tuple/list 차이를 없애 메모리의 계약과 디스크의 계약을
1469
+ 같은 표현으로 비교한다. 기존 수집 흔적이 있으면 누락된 계약을 새로
1470
+ 만들지 않는다. 그러면 과거 데이터의 의미를 새 계약으로 덮어쓰게 된다.
1471
+ """
1472
+ path = self.directory / 'contract.json'
1473
+ normalized = json.loads(json.dumps(contract, sort_keys=True))
1474
+ if path.exists():
1475
+ if json.loads(path.read_text()) != normalized:
1476
+ raise ValueError('custom collection provenance/selector/range/schema changed')
1477
+ else:
1478
+ if (any((self.directory / 'commits').glob('*.json'))
1479
+ or any((self.root / 'custom').glob('*/attempts/*'))):
1480
+ raise ValueError('incomplete custom collection manifest: missing contract')
1481
+ _atomic_json(path, normalized)
1482
+
1483
+ @staticmethod
1484
+ def key_id(key):
1485
+ """문서·choice·pass 키를 재시작 후에도 같은 marker 이름으로 변환한다."""
1486
+ import hashlib
1487
+ return hashlib.sha256(json.dumps(list(key), ensure_ascii=True).encode()).hexdigest()
1488
+
1489
+ def commit(self, key, attempt, rows_by_table, specs, indexes, coverage):
1490
+ """scalar 저장 → 전체 파일 hash → 완료 marker 순서로 한 수집 단위를 공개한다.
1491
+
1492
+ marker 이전에 실패한 파일은 미완료 attempt로 남아 기본 reader에서 제외된다.
1493
+ marker 이후에는 참조 파일 전체가 복구 가능해야 한다. fault checkpoint와
1494
+ 메모리의 completed 갱신도 이 공개 순서를 따른다.
1495
+ """
1496
+ import hashlib
1497
+
1498
+ tables = self._write_scalar_tables(attempt, rows_by_table, specs)
1499
+ # marker는 index 자체와 그 index가 가리키는 tensor를 모두 hash로 묶는다.
1500
+ files = list(tables.values()) + list(indexes)
1501
+ for index in indexes:
1502
+ files.append(json.loads((self.root / index).read_text())['path'])
1503
+ hashes = {p: hashlib.sha256((self.root / p).read_bytes()).hexdigest() for p in files}
1504
+ marker = {'version': 1, 'key': list(key), 'attempt': attempt, 'tables': tables,
1505
+ 'indexes': list(indexes), 'files': hashes, 'hooks': coverage}
1506
+ path = self.directory / 'commits' / (self.key_id(key) + '.json')
1507
+ if path.exists():
1508
+ raise ValueError('custom collection unit already committed')
1509
+ _custom_checkpoint("marker_before")
1510
+ _atomic_json(path, marker)
1511
+ _custom_checkpoint("marker_after")
1512
+ self.completed[tuple(key)] = marker
1513
+
1514
+ def _write_scalar_tables(
1515
+ self, attempt: str, rows_by_table: dict[str, list[dict[str, Any]]], specs: Sequence[Any],
1516
+ ) -> dict[str, str]:
1517
+ """attempt별 scalar 파일을 쓰고 run 디렉터리 기준 상대 경로를 반환한다.
1518
+
1519
+ 행이 없어도 스키마를 가진 파일을 쓴다. 유효한 빈 관측과 파일 누락을
1520
+ 구분하기 위해서다. 파일 fsync → rename → 디렉터리 fsync 순서를 유지하며,
1521
+ 여기서는 완료 marker를 쓰지 않아 아직 수집 결과로 공개되지 않는다.
1522
+ """
1523
+ tables = {}
1524
+ for spec in specs:
1525
+ name = f'custom/{spec.name}'
1526
+ directory = self.root / name / 'attempts' / attempt
1527
+ directory.mkdir(parents=True, exist_ok=True)
1528
+ path = directory / 'scalars.parquet'
1529
+ temporary = directory / 'scalars.tmp'
1530
+ rows = rows_by_table.get(name, [])
1531
+ schema = spec.table_spec().arrow_schema()
1532
+ table = pa.Table.from_pylist(rows, schema=schema)
1533
+ _custom_checkpoint("scalar_before")
1534
+ pq.write_table(table, temporary, compression='zstd')
1535
+ with temporary.open('rb') as stream:
1536
+ os.fsync(stream.fileno())
1537
+ os.replace(temporary, path)
1538
+ _sync_directory(path.parent)
1539
+ _custom_checkpoint("scalar_after")
1540
+ tables[name] = str(path.relative_to(self.root))
1541
+ return tables
1542
+
1543
+ def close(self):
1544
+ if self.lock is not None:
1545
+ self.lock.close()
1546
+ self.lock = None