codemaivanngu commited on
Commit
6b5bc0e
·
verified ·
1 Parent(s): ade3385

add read-only capture diagnostic for the 4096-cap parity bug

Browse files
Files changed (1) hide show
  1. diag/diagnose_parity_capture.py +78 -0
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)]