add read-only capture diagnostic for the 4096-cap parity bug
Browse files
diag/diagnose_parity_capture.py
CHANGED
|
@@ -143,6 +143,7 @@ def stored_pair_analysis(ids, mask, lengths, behavior_all, trainer_flat, band: i
|
|
| 143 |
positions: list[int] = []
|
| 144 |
rows: list[int] = []
|
| 145 |
cursor = 0
|
|
|
|
| 146 |
for r in range(ids.shape[0]):
|
| 147 |
n = int(lengths[r])
|
| 148 |
for p in mask[r, :n].nonzero().squeeze(-1).tolist():
|
|
@@ -152,6 +153,10 @@ def stored_pair_analysis(ids, mask, lengths, behavior_all, trainer_flat, band: i
|
|
| 152 |
b = float(behavior_all[r, p])
|
| 153 |
cursor += 1
|
| 154 |
if not math.isfinite(b):
|
|
|
|
|
|
|
|
|
|
|
|
|
| 155 |
continue
|
| 156 |
tv.append(t)
|
| 157 |
bv.append(b)
|
|
@@ -161,6 +166,8 @@ def stored_pair_analysis(ids, mask, lengths, behavior_all, trainer_flat, band: i
|
|
| 161 |
out: dict[str, Any] = {
|
| 162 |
"note": "stored behaviour vs stored trainer logprobs; this is the guard's delta",
|
| 163 |
"compared_tokens": len(delta),
|
|
|
|
|
|
|
| 164 |
"guard_stored_pair": stats(delta, positions),
|
| 165 |
"per_row": {},
|
| 166 |
"per_band": bands(delta, band),
|
|
@@ -294,6 +301,9 @@ def main() -> None:
|
|
| 294 |
print(f" guard delta from the capture itself: n={row['tokens']} "
|
| 295 |
f"mean={row['mean']:.6f} p99={row['p99']:.4f} max={row['max']:.4f} "
|
| 296 |
f">0.5={row['above_0p5']}")
|
|
|
|
|
|
|
|
|
|
| 297 |
if row.get("above_0p5_position_head"):
|
| 298 |
pos = row["above_0p5_position_head"]
|
| 299 |
print(f" first outlier positions: {pos[:24]}")
|
|
@@ -408,6 +418,74 @@ def main() -> None:
|
|
| 408 |
rows = [flat_row[i] for i in real]
|
| 409 |
|
| 410 |
d_T = [abs(a - b) for a, b in zip(fresh_T, behavior)]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 411 |
d_raw = [abs(a - b) for a, b in zip(fresh_raw, behavior)]
|
| 412 |
d_trainer = [abs(a - b) for a, b in zip(fresh_T, trainer)]
|
| 413 |
d_trainer_raw = [abs(a - b) for a, b in zip(fresh_raw, trainer)]
|
|
|
|
| 143 |
positions: list[int] = []
|
| 144 |
rows: list[int] = []
|
| 145 |
cursor = 0
|
| 146 |
+
nan_per_row: dict[str, int] = {}
|
| 147 |
for r in range(ids.shape[0]):
|
| 148 |
n = int(lengths[r])
|
| 149 |
for p in mask[r, :n].nonzero().squeeze(-1).tolist():
|
|
|
|
| 153 |
b = float(behavior_all[r, p])
|
| 154 |
cursor += 1
|
| 155 |
if not math.isfinite(b):
|
| 156 |
+
# NaN marks the synthetic terminal event on a capped row; it is
|
| 157 |
+
# intentional and carries no behaviour probability. Count it so the
|
| 158 |
+
# number can be checked against the capture's own per_sample stats.
|
| 159 |
+
nan_per_row[str(r)] = nan_per_row.get(str(r), 0) + 1
|
| 160 |
continue
|
| 161 |
tv.append(t)
|
| 162 |
bv.append(b)
|
|
|
|
| 166 |
out: dict[str, Any] = {
|
| 167 |
"note": "stored behaviour vs stored trainer logprobs; this is the guard's delta",
|
| 168 |
"compared_tokens": len(delta),
|
| 169 |
+
"nan_sentinel_tokens": sum(nan_per_row.values()),
|
| 170 |
+
"nan_sentinel_per_row": nan_per_row,
|
| 171 |
"guard_stored_pair": stats(delta, positions),
|
| 172 |
"per_row": {},
|
| 173 |
"per_band": bands(delta, band),
|
|
|
|
| 301 |
print(f" guard delta from the capture itself: n={row['tokens']} "
|
| 302 |
f"mean={row['mean']:.6f} p99={row['p99']:.4f} max={row['max']:.4f} "
|
| 303 |
f">0.5={row['above_0p5']}")
|
| 304 |
+
print(f" NaN terminal sentinels skipped: {stored['nan_sentinel_tokens']} "
|
| 305 |
+
f"per row {json.dumps(stored['nan_sentinel_per_row'])}"
|
| 306 |
+
" (1 per capped row is expected)")
|
| 307 |
if row.get("above_0p5_position_head"):
|
| 308 |
pos = row["above_0p5_position_head"]
|
| 309 |
print(f" first outlier positions: {pos[:24]}")
|
|
|
|
| 418 |
rows = [flat_row[i] for i in real]
|
| 419 |
|
| 420 |
d_T = [abs(a - b) for a, b in zip(fresh_T, behavior)]
|
| 421 |
+
|
| 422 |
+
# ---------------------------------------------------------------------
|
| 423 |
+
# Padded-batch recomputation, replicating `replay_parity.py --layout batch`.
|
| 424 |
+
#
|
| 425 |
+
# On the real capture the two recomputations disagree about which stored array
|
| 426 |
+
# they reproduce, so this has to be measured rather than argued: a per-row
|
| 427 |
+
# forward matches `stu_behavior_log_probs` while a padded-batch forward is what
|
| 428 |
+
# the trainer's own path used. Running BOTH in one process decides which side
|
| 429 |
+
# is anomalous, and whether padding is the mechanism at all.
|
| 430 |
+
# ---------------------------------------------------------------------
|
| 431 |
+
with torch.inference_mode():
|
| 432 |
+
pos_p = attn.long().cumsum(-1) - 1
|
| 433 |
+
pos_p.masked_fill_(attn == 0, 1)
|
| 434 |
+
logits_p = model(input_ids=ids.cuda(), attention_mask=attn.cuda(),
|
| 435 |
+
position_ids=pos_p.cuda(), use_cache=False).logits
|
| 436 |
+
padded_T_all: list[float] = []
|
| 437 |
+
padded_raw_all: list[float] = []
|
| 438 |
+
for r in range(ids.shape[0]):
|
| 439 |
+
n = int(lengths[r])
|
| 440 |
+
for p in mask[r, :n].nonzero().squeeze(-1).tolist():
|
| 441 |
+
j = p
|
| 442 |
+
if j < 0 or j >= n - 1:
|
| 443 |
+
padded_T_all.append(float("nan"))
|
| 444 |
+
padded_raw_all.append(float("nan"))
|
| 445 |
+
continue
|
| 446 |
+
block = logits_p[r, j].float()
|
| 447 |
+
label = int(ids[r, j + 1].item())
|
| 448 |
+
padded_T_all.append(float(torch.log_softmax(block / temperature, -1)[label]))
|
| 449 |
+
padded_raw_all.append(float(torch.log_softmax(block, -1)[label]))
|
| 450 |
+
del logits_p
|
| 451 |
+
torch.cuda.empty_cache()
|
| 452 |
+
padded_T = sub(padded_T_all)
|
| 453 |
+
padded_raw = sub(padded_raw_all)
|
| 454 |
+
report["padded_vs_row_fresh_T"] = stats([abs(a - b) for a, b in zip(padded_T, fresh_T)])
|
| 455 |
+
report["padded_fresh_T_vs_behavior"] = stats([abs(a - b) for a, b in zip(padded_T, behavior)])
|
| 456 |
+
report["padded_fresh_T_vs_stored_trainer"] = stats(
|
| 457 |
+
[abs(a - b) for a, b in zip(padded_T, trainer)])
|
| 458 |
+
report["padded_fresh_raw_vs_stored_trainer"] = stats(
|
| 459 |
+
[abs(a - b) for a, b in zip(padded_raw, trainer)])
|
| 460 |
+
report["row_fresh_raw_vs_stored_trainer"] = stats(
|
| 461 |
+
[abs(a - b) for a, b in zip(fresh_raw, trainer)])
|
| 462 |
+
report["row_fresh_raw_vs_behavior"] = stats(
|
| 463 |
+
[abs(a - b) for a, b in zip(fresh_raw, behavior)])
|
| 464 |
+
print("\n== which stored array does each recomputation reproduce? ==")
|
| 465 |
+
for key in ("padded_vs_row_fresh_T", "padded_fresh_T_vs_behavior",
|
| 466 |
+
"padded_fresh_T_vs_stored_trainer", "row_fresh_T_vs_behavior",
|
| 467 |
+
"row_fresh_T_vs_stored_trainer"):
|
| 468 |
+
if key in report:
|
| 469 |
+
r_ = report[key]
|
| 470 |
+
elif key == "row_fresh_T_vs_behavior":
|
| 471 |
+
r_ = report["guard_fresh_T_vs_behavior"]
|
| 472 |
+
else:
|
| 473 |
+
r_ = report["trainer_fresh_T_vs_stored_trainer"]
|
| 474 |
+
print(f" {key:36s} mean={r_['mean']:.6f} p99={r_['p99']:.4f} "
|
| 475 |
+
f"max={r_['max']:.4f} >0.5={r_['above_0p5']}")
|
| 476 |
+
print(" temperature probe on the TRAINER side (raw reference):")
|
| 477 |
+
for key in ("row_fresh_raw_vs_stored_trainer", "padded_fresh_raw_vs_stored_trainer"):
|
| 478 |
+
r_ = report[key]
|
| 479 |
+
print(f" {key:34s} mean={r_['mean']:.6f} p99={r_['p99']:.4f} "
|
| 480 |
+
f"max={r_['max']:.4f} >0.5={r_['above_0p5']}")
|
| 481 |
+
close = {"padded": report["padded_fresh_T_vs_stored_trainer"]["max"],
|
| 482 |
+
"row": report["trainer_fresh_T_vs_stored_trainer"]["max"]}
|
| 483 |
+
report["which_recomputation_matches_the_trainer_array"] = (
|
| 484 |
+
"padded_batch" if close["padded"] < close["row"] / 3 else
|
| 485 |
+
"per_row" if close["row"] < close["padded"] / 3 else "neither_clearly")
|
| 486 |
+
print(f" => the trainer's stored array is reproduced by: "
|
| 487 |
+
f"{report['which_recomputation_matches_the_trainer_array']}"
|
| 488 |
+
f" (padded max={close['padded']:.4f}, row max={close['row']:.4f})")
|
| 489 |
d_raw = [abs(a - b) for a, b in zip(fresh_raw, behavior)]
|
| 490 |
d_trainer = [abs(a - b) for a, b in zip(fresh_T, trainer)]
|
| 491 |
d_trainer_raw = [abs(a - b) for a, b in zip(fresh_raw, trainer)]
|