| """Evaluate a trained gen-time probe on the validation split. |
| |
| Mirrors ``train_probe_gen.py`` but runs in eval mode: loads the frozen base |
| LLaVA, the frozen SAE, and an existing probe checkpoint, then iterates the |
| val split, computes BCE / accuracy / AUROC per layer (and averaged), and |
| prints a summary. |
| |
| Usage (single-GPU): |
| python -m experiment.training.eval_probe_gen \ |
| --config experiment/adv_config.json \ |
| --probe_checkpoint /path/to/probes_gen_<rel>.pt |
| |
| Multi-GPU: |
| torchrun --nproc_per_node=4 -m experiment.training.eval_probe_gen \ |
| --config ... --probe_checkpoint ... |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| from datetime import datetime |
|
|
| import torch |
| import torch.distributed as dist |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, DistributedSampler |
| from transformers import AutoProcessor, AutoModelForPreTraining |
| from tqdm import tqdm |
|
|
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../..")) |
|
|
| from sae.Training_Utils import str_to_torch_dtype |
| from experiment.config.train_config import TrainConfig |
| from experiment.config.relation_config import get_relation_config |
| from experiment.data.datasets import FinetuneDataset, finetune_dataset_extra_kwargs |
| from experiment.training.finetune_adv import ( |
| FrozenSAEEncoder, |
| LayerProbes, |
| HiddenStateCapture, |
| count_lm_layers, |
| probe_labels, |
| layer_probes_from_checkpoint, |
| parse_args, |
| ) |
| from experiment.training.gen_features import ( |
| build_position_masks, |
| left_pad_collate, |
| masked_max_pool, |
| masked_mean_pool, |
| ) |
| from experiment.training.train_probe_gen import _generate_captions, _gen_features |
|
|
|
|
| def parse_extra_args_and_strip(): |
| parser = argparse.ArgumentParser(add_help=False) |
| parser.add_argument("--probe_checkpoint", type=str, required=True, |
| help="Path to probes_gen_*.pt produced by train_probe_gen.py") |
| parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) |
| parser.add_argument("--pool", choices=["max", "mean"], default="max", |
| help="Pooling over generated positions; must match how the probe was trained.") |
| parser.add_argument("--max_eval_samples", type=int, default=None, |
| help="Cap the number of val samples (overrides config.max_train_samples)") |
| parser.add_argument("--output_json", type=str, default=None, |
| help="Optional path to dump per-layer metrics as JSON") |
| extra, remaining = parser.parse_known_args() |
| sys.argv = [sys.argv[0]] + remaining |
| return extra |
|
|
|
|
| def _all_reduce_sum(t: torch.Tensor) -> torch.Tensor: |
| if dist.is_available() and dist.is_initialized(): |
| dist.all_reduce(t, op=dist.ReduceOp.SUM) |
| return t |
|
|
|
|
| def _gather_concat(local: torch.Tensor) -> torch.Tensor: |
| """Gather variable-length 1-D tensors across ranks; returns concatenated CPU tensor on rank 0.""" |
| if not (dist.is_available() and dist.is_initialized()): |
| return local.detach().cpu() |
| world = dist.get_world_size() |
| n_local = torch.tensor([local.numel()], device=local.device) |
| sizes = [torch.zeros_like(n_local) for _ in range(world)] |
| dist.all_gather(sizes, n_local) |
| max_n = int(max(s.item() for s in sizes)) |
| padded = torch.zeros(max_n, dtype=local.dtype, device=local.device) |
| padded[: local.numel()] = local |
| bufs = [torch.zeros_like(padded) for _ in range(world)] |
| dist.all_gather(bufs, padded) |
| chunks = [bufs[r][: int(sizes[r].item())] for r in range(world)] |
| return torch.cat(chunks, dim=0).detach().cpu() |
|
|
|
|
| def _auroc(scores: torch.Tensor, labels: torch.Tensor) -> float: |
| """Tie-aware AUROC computed via the Mann–Whitney U identity. Returns NaN if degenerate.""" |
| s = scores.float() |
| y = labels.float() |
| n_pos = int((y > 0.5).sum().item()) |
| n_neg = int((y < 0.5).sum().item()) |
| if n_pos == 0 or n_neg == 0: |
| return float("nan") |
| order = torch.argsort(s) |
| s_sorted = s[order] |
| y_sorted = y[order] |
| |
| n = s_sorted.numel() |
| ranks = torch.empty(n, dtype=torch.float64) |
| i = 0 |
| while i < n: |
| j = i |
| while j + 1 < n and s_sorted[j + 1] == s_sorted[i]: |
| j += 1 |
| avg = (i + j) / 2.0 + 1.0 |
| ranks[i : j + 1] = avg |
| i = j + 1 |
| sum_ranks_pos = ranks[(y_sorted > 0.5)].sum().item() |
| auroc = (sum_ranks_pos - n_pos * (n_pos + 1) / 2.0) / (n_pos * n_neg) |
| return float(auroc) |
|
|
|
|
| def main(): |
| extra = parse_extra_args_and_strip() |
| args, overrides = parse_args() |
|
|
| config = TrainConfig.load(args.config) |
| if args.relation: |
| config.relation = args.relation |
| if overrides: |
| config.apply_overrides(overrides) |
| config.resolve_from_relation() |
|
|
| relation_config = get_relation_config(config.relation) |
| adv_cfg = config.adv |
| model_dtype = str_to_torch_dtype(config.dtype) |
| assert adv_cfg.sae_checkpoint, "adv.sae_checkpoint must be set in config" |
| assert os.path.isfile(extra.probe_checkpoint), f"missing probe ckpt: {extra.probe_checkpoint}" |
|
|
| use_ddp = "LOCAL_RANK" in os.environ |
| if use_ddp: |
| local_rank = int(os.environ["LOCAL_RANK"]) |
| dist.init_process_group(backend="nccl") |
| torch.cuda.set_device(local_rank) |
| device = torch.device(f"cuda:{local_rank}") |
| is_main = local_rank == 0 |
| else: |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| local_rank = 0 |
| is_main = True |
|
|
| pool_fn = masked_mean_pool if extra.pool == "mean" else masked_max_pool |
|
|
| if is_main: |
| print(f"[eval_probe_gen] relation={config.relation} split={extra.split}") |
| print(f"[eval_probe_gen] sae={adv_cfg.sae_checkpoint}") |
| print(f"[eval_probe_gen] probe_ckpt={extra.probe_checkpoint}") |
| print(f"[eval_probe_gen] probe_label_mode={adv_cfg.probe_label_mode}") |
| print(f"[eval_probe_gen] pool={extra.pool}") |
|
|
| |
| model = AutoModelForPreTraining.from_pretrained( |
| config.model_name, torch_dtype=model_dtype, device_map={"": device}, |
| ) |
| processor = AutoProcessor.from_pretrained(config.model_name) |
| pad_token_id = processor.tokenizer.pad_token_id |
| if pad_token_id is None: |
| pad_token_id = processor.tokenizer.eos_token_id |
| for p in model.parameters(): |
| p.requires_grad_(False) |
| model.eval() |
| raw_model = model |
|
|
| n_layers = count_lm_layers(model) |
| probe_layers = adv_cfg.probe_layers if adv_cfg.probe_layers else list(range(n_layers)) |
|
|
| sae = FrozenSAEEncoder.from_checkpoint(adv_cfg.sae_checkpoint, device) |
| d_sae = sae.encoder.weight.shape[0] |
|
|
| probes = layer_probes_from_checkpoint( |
| extra.probe_checkpoint, probe_layers, d_sae, device=device |
| ) |
| probes.eval() |
| for p in probes.parameters(): |
| p.requires_grad_(False) |
| if is_main: |
| print(f"[eval_probe_gen] loaded probe with {len(probe_layers)} layers, d_sae={d_sae}") |
|
|
| max_samples = extra.max_eval_samples or config.max_train_samples |
| dataset = FinetuneDataset( |
| processor=processor, |
| prompt_config=config.prompts, |
| dataset_id=config.dataset_id, |
| scene_col=relation_config.scene_key, |
| object_col=relation_config.object_key, |
| csv_path=config.csv_path, |
| image_dir=config.image_dir, |
| max_samples=max_samples, |
| split=extra.split, |
| upsample_categories=None, |
| **finetune_dataset_extra_kwargs(config), |
| ) |
| if use_ddp: |
| sampler = DistributedSampler(dataset, shuffle=False, drop_last=False) |
| dataloader = DataLoader( |
| dataset, batch_size=config.batch_size, sampler=sampler, |
| num_workers=config.num_workers, pin_memory=True, drop_last=False, |
| ) |
| else: |
| dataloader = DataLoader( |
| dataset, batch_size=config.batch_size, shuffle=False, |
| num_workers=config.num_workers, pin_memory=True, drop_last=False, |
| ) |
|
|
| capture = HiddenStateCapture(raw_model, probe_layers) |
| n_image_patches = None |
|
|
| |
| L = len(probe_layers) |
| bce_sum = torch.zeros(L, dtype=torch.float64, device=device) |
| correct = torch.zeros(L, dtype=torch.float64, device=device) |
| n_total = torch.zeros((), dtype=torch.float64, device=device) |
| n_pos = torch.zeros((), dtype=torch.float64, device=device) |
| n_neg = torch.zeros((), dtype=torch.float64, device=device) |
| correct_pos = torch.zeros(L, dtype=torch.float64, device=device) |
| correct_neg = torch.zeros(L, dtype=torch.float64, device=device) |
| |
| all_probs_per_layer = [[] for _ in range(L)] |
| all_labels_local: list[torch.Tensor] = [] |
|
|
| pbar = tqdm(dataloader, desc=f"eval[{extra.split}]", disable=not is_main) |
| with torch.no_grad(): |
| for batch in pbar: |
| has_object = batch.pop("has_object").to(device) |
| is_scene = batch.pop("is_scene").to(device) |
| batch = {k: v.to(device) for k, v in batch.items()} |
| y = probe_labels(is_scene, has_object, adv_cfg.probe_label_mode) |
|
|
| _, full_seqs, prompt_lens, gen_lens = _generate_captions( |
| raw_model, |
| pixel_values=batch["pixel_values"], |
| input_ids=batch["input_ids"], |
| attention_mask=batch["attention_mask"], |
| max_new_tokens=adv_cfg.max_new_tokens_train, |
| do_sample=adv_cfg.gen_do_sample, |
| temperature=adv_cfg.gen_temperature, |
| pad_token_id=pad_token_id, |
| ) |
|
|
| tf_ids, tf_attn = left_pad_collate(full_seqs, pad_id=pad_token_id) |
| tf_ids = tf_ids.to(device) |
| tf_attn = tf_attn.to(device) |
|
|
| with capture: |
| raw_model( |
| pixel_values=batch["pixel_values"], |
| input_ids=tf_ids, |
| attention_mask=tf_attn, |
| use_cache=False, |
| ) |
|
|
| S = next(iter(capture.hidden_states.values())).shape[1] |
| if n_image_patches is None: |
| L_max = tf_ids.shape[1] |
| n_image_patches = S - L_max + 1 |
| if is_main: |
| print(f"[eval_probe_gen] inferred n_image_patches={n_image_patches} (S={S}, L_max={L_max})") |
|
|
| prompt_real_lens_t = torch.tensor(prompt_lens, device=device) |
| gen_lens_t = torch.tensor(gen_lens, device=device) |
| _, gen_mask = build_position_masks( |
| prompt_real_lens_t, gen_lens_t, n_image_patches, S |
| ) |
| keep = gen_lens_t > 0 |
| if not keep.any(): |
| continue |
|
|
| features = _gen_features(capture.hidden_states, sae, gen_mask, pool_fn) |
| features = {l: f[keep] for l, f in features.items()} |
| yk = y[keep] |
|
|
| logits_list = probes.forward_logits(features) |
| for i, z in enumerate(logits_list): |
| z32 = z.float() |
| bce = F.binary_cross_entropy_with_logits(z32, yk.float(), reduction="sum") |
| bce_sum[i] += bce.detach().double() |
| pred = (torch.sigmoid(z32) > 0.5).float() |
| correct[i] += (pred == yk.float()).sum().double() |
| correct_pos[i] += ((pred == 1) & (yk > 0.5)).sum().double() |
| correct_neg[i] += ((pred == 0) & (yk < 0.5)).sum().double() |
| all_probs_per_layer[i].append(torch.sigmoid(z32).detach()) |
|
|
| n_total += float(yk.numel()) |
| n_pos += float((yk > 0.5).sum().item()) |
| n_neg += float((yk < 0.5).sum().item()) |
| all_labels_local.append(yk.detach()) |
|
|
| if is_main: |
| pbar.set_postfix({ |
| "n": int(n_total.item()), |
| "p": int(n_pos.item()), |
| "n0": int(n_neg.item()), |
| }) |
|
|
| |
| _all_reduce_sum(bce_sum) |
| _all_reduce_sum(correct) |
| _all_reduce_sum(correct_pos) |
| _all_reduce_sum(correct_neg) |
| _all_reduce_sum(n_total) |
| _all_reduce_sum(n_pos) |
| _all_reduce_sum(n_neg) |
|
|
| |
| labels_local_cat = torch.cat(all_labels_local) if all_labels_local else torch.zeros(0, device=device) |
| labels_all = _gather_concat(labels_local_cat) |
| auroc_per_layer = [] |
| for i in range(L): |
| probs_local_cat = torch.cat(all_probs_per_layer[i]) if all_probs_per_layer[i] else torch.zeros(0, device=device) |
| probs_all = _gather_concat(probs_local_cat) |
| if is_main: |
| auroc_per_layer.append(_auroc(probs_all, labels_all)) |
|
|
| if not is_main: |
| if use_ddp: |
| dist.barrier() |
| dist.destroy_process_group() |
| return |
|
|
| n_total_v = float(n_total.item()) |
| n_pos_v = float(n_pos.item()) |
| n_neg_v = float(n_neg.item()) |
| bce_per_layer = (bce_sum / max(n_total_v, 1.0)).tolist() |
| acc_per_layer = (correct / max(n_total_v, 1.0)).tolist() |
| acc_pos_per_layer = (correct_pos / max(n_pos_v, 1.0)).tolist() if n_pos_v > 0 else [float("nan")] * L |
| acc_neg_per_layer = (correct_neg / max(n_neg_v, 1.0)).tolist() if n_neg_v > 0 else [float("nan")] * L |
|
|
| avg_bce = sum(bce_per_layer) / L |
| avg_acc = sum(acc_per_layer) / L |
| avg_acc_pos = sum(a for a in acc_pos_per_layer if a == a) / max(sum(1 for a in acc_pos_per_layer if a == a), 1) |
| avg_acc_neg = sum(a for a in acc_neg_per_layer if a == a) / max(sum(1 for a in acc_neg_per_layer if a == a), 1) |
| valid_aurocs = [a for a in auroc_per_layer if a == a] |
| avg_auroc = sum(valid_aurocs) / len(valid_aurocs) if valid_aurocs else float("nan") |
|
|
| print() |
| print("=" * 70) |
| print(f"[eval_probe_gen] results on split={extra.split}") |
| print("=" * 70) |
| print(f" N={int(n_total_v)} pos={int(n_pos_v)} neg={int(n_neg_v)} prevalence={n_pos_v / max(n_total_v, 1.0):.3f}") |
| print(f" AVG BCE : {avg_bce:.4f} (random≈{0.6931:.4f})") |
| print(f" AVG accuracy : {avg_acc:.4f}") |
| print(f" AVG acc | y=1 : {avg_acc_pos:.4f}") |
| print(f" AVG acc | y=0 : {avg_acc_neg:.4f}") |
| print(f" AVG AUROC : {avg_auroc:.4f}") |
| print() |
| print(f" {'layer':>6} | {'BCE':>8} | {'acc':>6} | {'acc|y=1':>8} | {'acc|y=0':>8} | {'AUROC':>6}") |
| print(f" {'-'*6} + {'-'*8} + {'-'*6} + {'-'*8} + {'-'*8} + {'-'*6}") |
| for i, l in enumerate(probe_layers): |
| au = auroc_per_layer[i] if i < len(auroc_per_layer) else float("nan") |
| print(f" {l:>6d} | {bce_per_layer[i]:>8.4f} | {acc_per_layer[i]:>6.4f} | " |
| f"{acc_pos_per_layer[i]:>8.4f} | {acc_neg_per_layer[i]:>8.4f} | {au:>6.4f}") |
|
|
| if extra.output_json: |
| os.makedirs(os.path.dirname(extra.output_json) or ".", exist_ok=True) |
| out = { |
| "split": extra.split, |
| "n_total": int(n_total_v), |
| "n_pos": int(n_pos_v), |
| "n_neg": int(n_neg_v), |
| "probe_layers": probe_layers, |
| "probe_label_mode": adv_cfg.probe_label_mode, |
| "pool": extra.pool, |
| "probe_checkpoint": extra.probe_checkpoint, |
| "averages": { |
| "bce": avg_bce, "accuracy": avg_acc, |
| "acc_pos": avg_acc_pos, "acc_neg": avg_acc_neg, |
| "auroc": avg_auroc, |
| }, |
| "per_layer": [ |
| { |
| "layer": int(l), |
| "bce": bce_per_layer[i], |
| "accuracy": acc_per_layer[i], |
| "acc_pos": acc_pos_per_layer[i], |
| "acc_neg": acc_neg_per_layer[i], |
| "auroc": auroc_per_layer[i], |
| } |
| for i, l in enumerate(probe_layers) |
| ], |
| "timestamp": datetime.now().isoformat(), |
| } |
| with open(extra.output_json, "w") as f: |
| json.dump(out, f, indent=2) |
| print(f"[eval_probe_gen] wrote {extra.output_json}") |
|
|
| if use_ddp: |
| dist.barrier() |
| dist.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|