Skip to content

Commit cb8902d

Browse files
Ronald Tseronaldtse
authored andcommitted
feat(parity): per-pair flip-position dump — the unit the flip bootstrap needs
run_margin_analysis(dump_positions=...) writes one JSONL record per pair (tokens, flip positions, their reference margins) beside the aggregate margins JSON. MarginReport stays schema-identical; the dump is diagnostic. Enables the TODO.training-work/05 shipped-vs-head32 flip CIs; aggregate flip counts in the dump must reconcile with the report (tested).
1 parent 3b39c21 commit cb8902d

2 files changed

Lines changed: 31 additions & 1 deletion

File tree

‎src/imf/parity.py‎

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -253,7 +253,7 @@ def _onnx_forced_logits(enc_sess, dec_sess, source: str, target_ids: list[int]):
253253
return dict(zip(out_names, out, strict=True))["logits"][0]
254254

255255

256-
def run_margin_analysis(model, zip_path, pairs, max_len: int = 256) -> MarginReport:
256+
def run_margin_analysis(model, zip_path, pairs, max_len: int = 256, dump_positions: Path | str | None = None) -> MarginReport:
257257
"""pairs: iterable of (source_text, gold_target) — the same probe set
258258
the CER parity gate uses. Teacher-forces both sides and measures the
259259
argmax flip rate, reference top1−top2 margin quantiles, and KL
@@ -286,6 +286,20 @@ def run_margin_analysis(model, zip_path, pairs, max_len: int = 256) -> MarginRep
286286

287287
margins = np.concatenate(margin_chunks)
288288
flips = np.concatenate(flip_chunks)
289+
if dump_positions is not None:
290+
# per-position flip vectors, pair-indexed — the unit the
291+
# paired bootstrap needs for shipped-vs-rebuild comparisons
292+
import json as _json
293+
294+
dump = Path(dump_positions)
295+
dump.parent.mkdir(parents=True, exist_ok=True)
296+
with dump.open("w", encoding="utf-8") as out:
297+
for pair_idx, (m, f) in enumerate(zip(margin_chunks, flip_chunks)):
298+
out.write(_json.dumps(
299+
{"pair": pair_idx, "tokens": int(f.size),
300+
"flip_positions": [int(x) for x in np.nonzero(f)[0]],
301+
"flip_margins": [round(float(m[x]), 4) for x in np.nonzero(f)[0]]}
302+
) + "\n")
289303
p1, p10, p50 = (float(np.quantile(margins, q)) for q in (0.01, 0.10, 0.50))
290304
n_flips = int(flips.sum())
291305
return MarginReport(

‎tests/test_imf_parity.py‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -166,3 +166,19 @@ def test_margin_report_json_roundtrip(gated_zip: Path, fixture_model, tmp_path:
166166
}
167167
assert data["samples"] == report.samples
168168
assert data["flip_rate"] == report.flip_rate
169+
170+
171+
def test_margin_analysis_dumps_per_pair_positions(gated_zip: Path, fixture_model, tmp_path: Path) -> None:
172+
"""TODO.training-work/05: the flip bootstrap needs per-pair token and
173+
flip counts, not just aggregates."""
174+
dump = tmp_path / "positions.jsonl"
175+
report = run_margin_analysis(
176+
fixture_model, gated_zip, PAIRS, max_len=12, dump_positions=dump,
177+
)
178+
rows = [json.loads(line) for line in dump.read_text().splitlines() if line]
179+
assert rows and len(rows) <= len(PAIRS)
180+
assert all({"pair", "tokens", "flip_positions", "flip_margins"} <= set(r) for r in rows)
181+
total_flips = sum(len(r["flip_positions"]) for r in rows)
182+
assert total_flips == report.flipped_tokens
183+
total_tokens = sum(r["tokens"] for r in rows)
184+
assert total_tokens == report.tokens

0 commit comments

Comments
 (0)