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

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

Browse files
Files changed (1) hide show
  1. 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]] = []