add read-only capture diagnostic for the 4096-cap parity bug
Browse files- diag/diagnose_parity_capture.py +135 -1
diag/diagnose_parity_capture.py
CHANGED
|
@@ -124,6 +124,79 @@ def repeat_fraction(ids: list[int], lag: int) -> float:
|
|
| 124 |
return sum(1 for i in range(len(ids) - lag) if ids[i] == ids[i + lag]) / (len(ids) - lag)
|
| 125 |
|
| 126 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
def score_row(model, row_ids: list[int], temperature: float) -> dict[str, list[float]]:
|
| 128 |
"""Per-position logprobs for positions 1..L-1 of one unpadded row."""
|
| 129 |
import torch
|
|
@@ -169,6 +242,9 @@ def main() -> None:
|
|
| 169 |
help="also rescore the capture's exact tokens with SGLang under the "
|
| 170 |
"production flags, to separate 'the stored values were corrupted' "
|
| 171 |
"from 'the engine is wrong today'")
|
|
|
|
|
|
|
|
|
|
| 172 |
args = ap.parse_args()
|
| 173 |
|
| 174 |
import torch
|
|
@@ -198,12 +274,70 @@ def main() -> None:
|
|
| 198 |
"attention_backend", "source_commit")}, sort_keys=True),
|
| 199 |
flush=True)
|
| 200 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 201 |
model = AutoModelForCausalLM.from_pretrained(
|
| 202 |
args.capture / "student", dtype=getattr(torch, args.dtype),
|
| 203 |
attn_implementation="eager", local_files_only=True,
|
| 204 |
trust_remote_code=False).cuda()
|
| 205 |
model.train(bool(meta.get("model_training")))
|
| 206 |
-
lengths = attn.sum(-1).tolist()
|
| 207 |
|
| 208 |
# ---- per-row fresh recomputation (unpadded) -------------------------------
|
| 209 |
per_row: list[dict[str, Any]] = []
|
|
|
|
| 124 |
return sum(1 for i in range(len(ids) - lag) if ids[i] == ids[i + lag]) / (len(ids) - lag)
|
| 125 |
|
| 126 |
|
| 127 |
+
def stored_pair_analysis(ids, mask, lengths, behavior_all, trainer_flat, band: int,
|
| 128 |
+
max_shift: int) -> dict[str, Any]:
|
| 129 |
+
"""Structure of the delta that actually tripped the guard, with NO model.
|
| 130 |
+
|
| 131 |
+
`batch.pt` carries both sides of the comparison: `stu_behavior_log_probs` (what
|
| 132 |
+
SGLang returned) and `actual_log_probs` (what the trainer computed, masked-flat
|
| 133 |
+
in row-major order). Their difference IS the guard's delta. That means the
|
| 134 |
+
structural questions -- which row, where, contiguous or strided, what period --
|
| 135 |
+
are answerable from the capture alone, with no checkpoint, no GPU and no HF.
|
| 136 |
+
|
| 137 |
+
This matters because the capture holds a 6.4 GB student checkpoint and the bug
|
| 138 |
+
note flags a cleanup policy for it: if it is gone, a model-dependent diagnostic
|
| 139 |
+
produces nothing at all, while this path still reports the primary evidence.
|
| 140 |
+
"""
|
| 141 |
+
tv: list[float] = [] # stored trainer logprobs, aligned order
|
| 142 |
+
bv: list[float] = [] # stored behaviour logprobs, aligned order
|
| 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():
|
| 149 |
+
if cursor >= len(trainer_flat):
|
| 150 |
+
break
|
| 151 |
+
t = float(trainer_flat[cursor])
|
| 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)
|
| 158 |
+
positions.append(p)
|
| 159 |
+
rows.append(r)
|
| 160 |
+
delta = [abs(a - b) for a, b in zip(tv, bv)]
|
| 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),
|
| 167 |
+
"shift_test_stored_pair": {},
|
| 168 |
+
"periodicity": {},
|
| 169 |
+
"outlier_token_id_histogram_top": {},
|
| 170 |
+
}
|
| 171 |
+
for r in range(ids.shape[0]):
|
| 172 |
+
sel = [i for i, rr in enumerate(rows) if rr == r]
|
| 173 |
+
if sel:
|
| 174 |
+
out["per_row"][str(r)] = stats([delta[i] for i in sel],
|
| 175 |
+
[positions[i] for i in sel])
|
| 176 |
+
# Alignment test between the two STORED arrays: compare trainer[i] against
|
| 177 |
+
# behaviour[i + k]. If the engine's values were shifted relative to the
|
| 178 |
+
# trainer's, a non-zero shift would beat shift 0 here. (Comparing the deltas to
|
| 179 |
+
# each other, as a first draft did, is meaningless.)
|
| 180 |
+
for k in range(-max_shift, max_shift + 1):
|
| 181 |
+
pairs = [abs(tv[i] - bv[i + k]) for i in range(len(tv)) if 0 <= i + k < len(bv)]
|
| 182 |
+
out["shift_test_stored_pair"][str(k)] = stats(pairs)
|
| 183 |
+
for r in range(ids.shape[0]):
|
| 184 |
+
n = int(lengths[r])
|
| 185 |
+
row_ids = [int(v) for v in ids[r, :n].tolist()]
|
| 186 |
+
out["periodicity"][str(r)] = {
|
| 187 |
+
str(lag): round(repeat_fraction(row_ids, lag), 4)
|
| 188 |
+
for lag in (1, 2, 3, 4, 5, 6, 7, 8, 12, 16, 24, 32)}
|
| 189 |
+
outliers = [i for i, v in enumerate(delta) if v > 0.5]
|
| 190 |
+
if outliers:
|
| 191 |
+
tids = [int(ids[rows[i], positions[i]].item()) for i in outliers]
|
| 192 |
+
hist: dict[str, int] = {}
|
| 193 |
+
for t in tids:
|
| 194 |
+
hist[str(t)] = hist.get(str(t), 0) + 1
|
| 195 |
+
out["outlier_token_id_histogram_top"] = dict(
|
| 196 |
+
sorted(hist.items(), key=lambda kv: -kv[1])[:10])
|
| 197 |
+
return out
|
| 198 |
+
|
| 199 |
+
|
| 200 |
def score_row(model, row_ids: list[int], temperature: float) -> dict[str, list[float]]:
|
| 201 |
"""Per-position logprobs for positions 1..L-1 of one unpadded row."""
|
| 202 |
import torch
|
|
|
|
| 242 |
help="also rescore the capture's exact tokens with SGLang under the "
|
| 243 |
"production flags, to separate 'the stored values were corrupted' "
|
| 244 |
"from 'the engine is wrong today'")
|
| 245 |
+
ap.add_argument("--no-model", action="store_true",
|
| 246 |
+
help="skip every checkpoint-dependent check and report only the "
|
| 247 |
+
"stored-pair structure (works from batch.pt alone)")
|
| 248 |
args = ap.parse_args()
|
| 249 |
|
| 250 |
import torch
|
|
|
|
| 274 |
"attention_backend", "source_commit")}, sort_keys=True),
|
| 275 |
flush=True)
|
| 276 |
|
| 277 |
+
lengths = attn.sum(-1).tolist()
|
| 278 |
+
|
| 279 |
+
# ---------------------------------------------------------------------
|
| 280 |
+
# Model-free structural analysis, ALWAYS run first.
|
| 281 |
+
#
|
| 282 |
+
# `batch.pt` holds both sides of the guard's comparison, so the structure of
|
| 283 |
+
# the failure is available with no checkpoint, no GPU and no HF. The capture
|
| 284 |
+
# carries a 6.4 GB student checkpoint and the bug note flags a cleanup policy
|
| 285 |
+
# for it; if it is gone, a model-dependent diagnostic reports nothing at all,
|
| 286 |
+
# while this path still reports the primary evidence.
|
| 287 |
+
# ---------------------------------------------------------------------
|
| 288 |
+
trainer_flat = batch["actual_log_probs"].tolist()
|
| 289 |
+
stored = stored_pair_analysis(ids, mask, lengths, behavior_all, trainer_flat,
|
| 290 |
+
args.band, args.max_shift)
|
| 291 |
+
report["stored_pair"] = stored
|
| 292 |
+
print("\n== STORED-PAIR structure (no model needed) ==")
|
| 293 |
+
row = stored["guard_stored_pair"]
|
| 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]}")
|
| 300 |
+
# stats() already computes the gap histogram over the first 300 outliers.
|
| 301 |
+
print(f" gap histogram : "
|
| 302 |
+
f"{json.dumps(row.get('above_0p5_gap_histogram') or {})}")
|
| 303 |
+
mod: dict[str, int] = {}
|
| 304 |
+
for p in pos:
|
| 305 |
+
mod[str(p % 8)] = mod.get(str(p % 8), 0) + 1
|
| 306 |
+
print(f" mod-8 histogram : {json.dumps(mod)}")
|
| 307 |
+
print(f" span : {row.get('above_0p5_span_positions')}")
|
| 308 |
+
for r, s in stored["per_row"].items():
|
| 309 |
+
print(f" row {r}: n={s['tokens']} mean={s['mean']:.6f} p99={s['p99']:.4f} "
|
| 310 |
+
f"max={s['max']:.4f} >0.5={s['above_0p5']}")
|
| 311 |
+
print(" per band (mean/max): " + json.dumps(
|
| 312 |
+
[[b[0], b[1], round(b[3], 5), round(b[4], 4)] for b in stored["per_band"]]))
|
| 313 |
+
print(" shift test (stored pair): " + json.dumps(
|
| 314 |
+
{k: [round(v["mean"], 6), round(v["max"], 4)]
|
| 315 |
+
for k, v in stored["shift_test_stored_pair"].items()}))
|
| 316 |
+
print(" periodicity: " + json.dumps(stored["periodicity"]))
|
| 317 |
+
if stored["outlier_token_id_histogram_top"]:
|
| 318 |
+
print(" outlier token-id histogram (top): "
|
| 319 |
+
+ json.dumps(stored["outlier_token_id_histogram_top"]))
|
| 320 |
+
|
| 321 |
+
have_model = not args.no_model and (args.capture / "student").is_dir()
|
| 322 |
+
if not have_model:
|
| 323 |
+
reason = ("--no-model was requested" if args.no_model
|
| 324 |
+
else f"no student checkpoint at {args.capture / 'student'}")
|
| 325 |
+
report["model_skipped"] = reason
|
| 326 |
+
print(f"\nPARITY_DIAG_MODEL=skipped ({reason})")
|
| 327 |
+
print(" The HF cross-checks (temperature detector, entropy/margin, padding,"
|
| 328 |
+
" engine) need the checkpoint and were not run.")
|
| 329 |
+
if args.out:
|
| 330 |
+
args.out.write_text(json.dumps(report, indent=2, sort_keys=True,
|
| 331 |
+
default=str) + "\n")
|
| 332 |
+
print("PARITY_DIAG_OUT=" + str(args.out), flush=True)
|
| 333 |
+
print("PARITY_DIAG_DONE=1", flush=True)
|
| 334 |
+
return
|
| 335 |
+
|
| 336 |
model = AutoModelForCausalLM.from_pretrained(
|
| 337 |
args.capture / "student", dtype=getattr(torch, args.dtype),
|
| 338 |
attn_implementation="eager", local_files_only=True,
|
| 339 |
trust_remote_code=False).cuda()
|
| 340 |
model.train(bool(meta.get("model_training")))
|
|
|
|
| 341 |
|
| 342 |
# ---- per-row fresh recomputation (unpadded) -------------------------------
|
| 343 |
per_row: list[dict[str, Any]] = []
|