Skip to content

Commit 3d7b934

Browse files
andre15silvaclaude
andcommitted
feat(probe): add --shuffle-labels flag for sanity-check ablation
Randomly permutes labels within each split (train/val/test independently) before training. Results stored under a separate run_id (e.g. *_shuffled) reusing the existing cache. AUC should be ~0.5 if the probe is not memorizing. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
1 parent 6fed20a commit 3d7b934

2 files changed

Lines changed: 14 additions & 0 deletions

File tree

run_probe.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,8 @@ def main():
2020
default="position",
2121
help="Axis to bin on during evaluation: token position, relative step, or exact step number")
2222
parser.add_argument("--probe-arch", choices=["linear", "mlp"], default="linear")
23+
parser.add_argument("--shuffle-labels", action="store_true",
24+
help="Randomly permute labels within each split before training (sanity-check baseline).")
2325

2426
subparsers = parser.add_subparsers(dest="mode", required=True)
2527

@@ -53,6 +55,7 @@ def main():
5355
probe_arch=args.probe_arch,
5456
n_eval_bins=args.n_eval_bins,
5557
eval_bin_axis=args.eval_bin_axis,
58+
shuffle_labels=args.shuffle_labels,
5659
)
5760
else:
5861
run_final(
@@ -71,6 +74,7 @@ def main():
7174
probe_arch=args.probe_arch,
7275
n_eval_bins=args.n_eval_bins,
7376
eval_bin_axis=args.eval_bin_axis,
77+
shuffle_labels=args.shuffle_labels,
7478
)
7579

7680

src/probe.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,7 @@ def train_probe_layer(
9696
log_fn: Callable[[dict], None] | None = None,
9797
n_eval_bins: int | None = None,
9898
eval_bin_axis: str = "position",
99+
shuffle_labels: bool = False,
99100
) -> list[ProbeResult]:
100101
_set_seeds(seed)
101102
data = torch.load(cache_path, weights_only=False)
@@ -156,6 +157,11 @@ def train_probe_layer(
156157
H_train, y_train = bin_H[train_mask], bin_y[train_mask]
157158
H_val, y_val = bin_H[val_mask], bin_y[val_mask]
158159
H_test, y_test = bin_H[test_mask], bin_y[test_mask]
160+
161+
if shuffle_labels:
162+
y_train = y_train[torch.randperm(len(y_train))]
163+
y_val = y_val[torch.randperm(len(y_val))]
164+
y_test = y_test[torch.randperm(len(y_test))]
159165
rel_pos_val = bin_rel_pos[val_mask]
160166
rel_pos_test = bin_rel_pos[test_mask]
161167
step_idx_val = bin_step_idx[val_mask] if bin_step_idx is not None else None
@@ -333,6 +339,7 @@ def run_sweep(
333339
probe_arch: str = "linear",
334340
n_eval_bins: int | None = None,
335341
eval_bin_axis: str = "position",
342+
shuffle_labels: bool = False,
336343
) -> None:
337344
import wandb
338345
cache_base = Path(cache_dir) / (cache_run_id or run_id) / probe_name
@@ -366,6 +373,7 @@ def log_fn(metrics: dict) -> None:
366373
batch_size=cfg.batch_size, patience=cfg.patience,
367374
seed=seed, n_bins=n_bins, probe_arch=probe_arch,
368375
log_fn=log_fn, n_eval_bins=n_eval_bins, eval_bin_axis=eval_bin_axis,
376+
shuffle_labels=shuffle_labels,
369377
)
370378
wandb.log({
371379
"mean_val_f1": np.mean([r.val_f1 for r in results]) if results else 0.0,
@@ -392,6 +400,7 @@ def run_final(
392400
probe_arch: str = "linear",
393401
n_eval_bins: int | None = None,
394402
eval_bin_axis: str = "position",
403+
shuffle_labels: bool = False,
395404
) -> dict:
396405
import wandb
397406
cache_base = Path(cache_dir) / (cache_run_id or run_id) / probe_name
@@ -425,6 +434,7 @@ def log_fn(metrics: dict) -> None:
425434
batch_size=batch_size, patience=patience,
426435
seed=seed, n_bins=n_bins, probe_arch=probe_arch,
427436
log_fn=make_log_fn(layer_idx), n_eval_bins=n_eval_bins, eval_bin_axis=eval_bin_axis,
437+
shuffle_labels=shuffle_labels,
428438
)
429439
all_results[layer_idx] = results
430440
for r in results:

0 commit comments

Comments
 (0)