Buckets:
| #!/usr/bin/env python3 | |
| """Task 4 follow-up: PARTIAL (quarter) ablations along token OR hidden axis. | |
| For the 180-sample focus set (9 classes), for every predictor layer we replace | |
| only ONE QUARTER of that layer's residual-stream activation with a substitute, | |
| leaving the other three quarters intact. Two axes, run back to back: | |
| * axis="token" -- split the 256 tokens into 4 groups of 64; replace one group. | |
| * axis="hidden" -- split the 384 hidden dims into 4 groups of 96; replace one. | |
| Modes per quarter: zero / mean (global dataset centroid) / resample (donor image | |
| i+1's activation). Because only a quarter is overwritten, the untouched state | |
| still carries the current image, so (unlike the full-tensor swap) results vary by | |
| layer and quarter and resample is no longer degenerate. | |
| """ | |
| # ruff: noqa: N803, N806 -- names mirror this analysis's linear-algebra notation | |
| # (e.g. L for layer, N for sample count) rather than PEP8 casing. | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import os | |
| import sys | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| import matplotlib | |
| import numpy as np | |
| import torch | |
| import yaml | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| TASK_ROOT = Path(__file__).resolve().parents[1] | |
| REPO = str(TASK_ROOT.parent) | |
| if REPO not in sys.path: | |
| sys.path.insert(0, REPO) | |
| from world_model_lens import LatentProber # noqa: E402 | |
| from world_model_lens.analysis.ood_detection import MahalanobisOODDetector # noqa: E402 | |
| from world_model_lens.data import load_imagenet_image # noqa: E402 | |
| from world_model_lens.hub.model_hub import ModelHub # noqa: E402 | |
| SCRATCH = Path(__file__).resolve().parent | |
| FOCUS = json.load(open(SCRATCH / "focus_180.json")) | |
| OUT_ROOT = Path( | |
| os.environ.get("WML_OUTPUT_ROOT", str(TASK_ROOT / "outputs" / "ijepa_task4_quarters")) | |
| ) | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32 | |
| BATCH = 32 | |
| SEED = 42 | |
| NQ = 4 # quarters | |
| MODES = ["zero", "mean", "resample"] | |
| TARGET_PATCH_SIDE = 4 | |
| PROBE_SPLIT = 0.2 | |
| # stable class ordering -> contiguous probe labels 0..8, keep folder as name | |
| FOLDERS = sorted({s["class_name"] for s in FOCUS}, key=lambda x: int(x)) | |
| FOLDER_TO_PLABEL = {f: i for i, f in enumerate(FOLDERS)} | |
| HUMAN = { | |
| "980": "volcano", | |
| "913": "shipwreck", | |
| "825": "stone wall", | |
| "754": "radio", | |
| "517": "crane", | |
| "225": "malinois", | |
| "250": "husky", | |
| "238": "swiss-mtn-dog", | |
| "348": "ram", | |
| } | |
| def seed_all(): | |
| import random | |
| random.seed(SEED) | |
| np.random.seed(SEED) | |
| torch.manual_seed(SEED) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(SEED) | |
| def load_model(): | |
| a = ModelHub.load("ijepa-vit-h-in1k", device=DEVICE) | |
| a = a.half() if DTYPE == torch.float16 else a | |
| a.eval() | |
| return a | |
| def fixed_masks(a): | |
| n = int(a.context_encoder.patch_embed.n_patches) | |
| grid = int(math.sqrt(n)) | |
| side = TARGET_PATCH_SIDE | |
| start = (grid - side) // 2 | |
| target = {r * grid + c for r in range(start, start + side) for c in range(start, start + side)} | |
| context = [p for p in range(n) if p not in target] | |
| return context, sorted(target), n | |
| def collect_clean(a, context_ids, target_ids, layers): | |
| store = {} | |
| def cap(L): | |
| def h(_m, _i, o): | |
| store[L] = o.detach() | |
| return h | |
| ctx_c, pred_c, tgt_c = [], [], [] | |
| act_c = {L: [] for L in layers} | |
| idxs = list(range(len(FOCUS))) | |
| for s in range(0, len(idxs), BATCH): | |
| batch = idxs[s : s + BATCH] | |
| tens = [load_imagenet_image(FOCUS[i]["path"], image_size=224) for i in batch] | |
| obs = torch.cat(tens, 0).to(DEVICE, DTYPE) | |
| handles = [ | |
| a.predictor.blocks[L].hook_resid_post.register_forward_hook(cap(L)) for L in layers | |
| ] | |
| try: | |
| cl = a.context_encoder(obs, patch_ids=context_ids) | |
| pred = a.predictor(cl, context_ids, target_ids) | |
| tgt = a.target_encoder(obs) | |
| finally: | |
| for hd in handles: | |
| hd.remove() | |
| ctx_c.append(cl.detach().to("cpu", torch.float16)) | |
| pred_c.append(pred.detach().to("cpu", torch.float16)) | |
| tgt_c.append(tgt[:, target_ids, :].detach().to("cpu", torch.float16)) | |
| for L in layers: | |
| act_c[L].append(store[L].to("cpu", torch.float16)) | |
| return { | |
| "context": torch.cat(ctx_c, 0), | |
| "pred": torch.cat(pred_c, 0), | |
| "tgt": torch.cat(tgt_c, 0), | |
| "act": {L: torch.cat(act_c[L], 0) for L in layers}, | |
| } | |
| def intervene(a, ctx_batch, context_ids, target_ids, L, hook): | |
| hd = a.predictor.blocks[L].hook_resid_post.register_forward_hook(hook) | |
| try: | |
| pred = a.predictor(ctx_batch, context_ids, target_ids) | |
| finally: | |
| hd.remove() | |
| return pred.detach() | |
| def quarter_slices(axis, n_tokens, n_hidden): | |
| if axis == "token": | |
| step = n_tokens // NQ | |
| return [slice(q * step, (q + 1) * step) for q in range(NQ)], ("token", n_tokens, step) | |
| step = n_hidden // NQ | |
| return [slice(q * step, (q + 1) * step) for q in range(NQ)], ("hidden", n_hidden, step) | |
| def make_hook(axis, sl, mode, mean_dev, donor_dev): | |
| def hook(_m, _i, out): | |
| o = out.clone() | |
| if axis == "token": | |
| if mode == "zero": | |
| o[:, sl, :] = 0 | |
| elif mode == "mean": | |
| o[:, sl, :] = mean_dev[:, sl, :].to(o.dtype) | |
| else: | |
| o[:, sl, :] = donor_dev[:, sl, :].to(o.dtype) | |
| else: | |
| if mode == "zero": | |
| o[:, :, sl] = 0 | |
| elif mode == "mean": | |
| o[:, :, sl] = mean_dev[:, :, sl].to(o.dtype) | |
| else: | |
| o[:, :, sl] = donor_dev[:, :, sl].to(o.dtype) | |
| return o | |
| return hook | |
| def metrics(pred, tgt, clean, mod_feat, ld, pd): | |
| p = pred.float().flatten() | |
| t = tgt.float().flatten() | |
| c = clean.float().flatten() | |
| mse = torch.mean((p - t) ** 2) | |
| cmse = torch.mean((c - t) ** 2) | |
| return { | |
| "prediction_mse": float(mse), | |
| "prediction_mse_ratio": float(mse / cmse.clamp_min(1e-12)), | |
| "target_cosine": float(torch.nn.functional.cosine_similarity(p, t, 0)), | |
| "prediction_shift_l2": float((p - c).norm()), | |
| "clean_prediction_cosine": float(torch.nn.functional.cosine_similarity(p, c, 0)), | |
| "substitution_maha": float(ld.score(mod_feat).item()), | |
| "prediction_maha": float(pd.score(pred.float().mean(0, keepdim=True)).item()), | |
| } | |
| def run_axis( | |
| a, axis, context_ids, target_ids, layers, clean, layer_means, ldet, pdet, n_tokens, n_hidden | |
| ): | |
| ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") | |
| run_dir = OUT_ROOT / axis / ts | |
| run_dir.mkdir(parents=True, exist_ok=True) | |
| slices, meta = quarter_slices(axis, n_tokens, n_hidden) | |
| N = len(FOCUS) | |
| labels = [FOLDER_TO_PLABEL[FOCUS[i]["class_name"]] for i in range(N)] | |
| rows = [] | |
| feat_sets = {} | |
| # clean baseline rows (one per layer for aggregation symmetry; layer-invariant) | |
| clean_feat = clean["pred"].float().mean(1).tolist() | |
| for L in layers: | |
| for i in range(N): | |
| p = clean["pred"][i].float().flatten() | |
| t = clean["tgt"][i].float().flatten() | |
| rows.append( | |
| { | |
| "sample_index": i, | |
| "label": labels[i], | |
| "class_name": FOCUS[i]["class_name"], | |
| "human": HUMAN[FOCUS[i]["class_name"]], | |
| "layer": L, | |
| "quarter": -1, | |
| "mode": "clean", | |
| "prediction_mse": float(torch.mean((p - t) ** 2)), | |
| "prediction_mse_ratio": 1.0, | |
| "target_cosine": float(torch.nn.functional.cosine_similarity(p, t, 0)), | |
| "prediction_shift_l2": 0.0, | |
| "clean_prediction_cosine": 1.0, | |
| "substitution_maha": None, | |
| "prediction_maha": None, | |
| } | |
| ) | |
| feat_sets[(L, -1, "clean")] = clean_feat | |
| batches = [list(range(s, min(s + BATCH, N))) for s in range(0, N, BATCH)] | |
| for L in layers: | |
| lm = layer_means[L] # [seq, hidden] float, cpu | |
| ld = ldet[L] | |
| clean_act_L = clean["act"][L].float() # [N, seq, hidden] cpu | |
| for q, sl in enumerate(slices): | |
| for mode in MODES: | |
| feats = [] | |
| for batch in batches: | |
| ctx = clean["context"][batch].to(DEVICE, DTYPE) | |
| B = len(batch) | |
| mean_dev = lm.unsqueeze(0).expand(B, -1, -1).to(DEVICE) | |
| if mode == "resample": | |
| donor_idx = [(i + 1) % N for i in batch] | |
| donor_dev = clean_act_L[donor_idx].to(DEVICE) | |
| else: | |
| donor_dev = None | |
| hook = make_hook(axis, sl, mode, mean_dev, donor_dev) | |
| pred = intervene(a, ctx, context_ids, target_ids, L, hook).to( | |
| "cpu", torch.float32 | |
| ) | |
| feats.extend(pred.mean(1).tolist()) | |
| # build modified full activation (for substitution_maha) on cpu | |
| mod = clean_act_L[batch].clone() | |
| if axis == "token": | |
| if mode == "zero": | |
| mod[:, sl, :] = 0 | |
| elif mode == "mean": | |
| mod[:, sl, :] = lm[sl, :].unsqueeze(0) | |
| else: | |
| mod[:, sl, :] = clean_act_L[donor_idx][:, sl, :] | |
| else: | |
| if mode == "zero": | |
| mod[:, :, sl] = 0 | |
| elif mode == "mean": | |
| mod[:, :, sl] = lm[:, sl].unsqueeze(0) | |
| else: | |
| mod[:, :, sl] = clean_act_L[donor_idx][:, :, sl] | |
| mod_feat = mod.mean(1) # [B, hidden] | |
| for pos, i in enumerate(batch): | |
| m = metrics( | |
| pred[pos], | |
| clean["tgt"][i].float(), | |
| clean["pred"][i].float(), | |
| mod_feat[pos : pos + 1], | |
| ld, | |
| pdet, | |
| ) | |
| rows.append( | |
| { | |
| "sample_index": i, | |
| "label": labels[i], | |
| "class_name": FOCUS[i]["class_name"], | |
| "human": HUMAN[FOCUS[i]["class_name"]], | |
| "layer": L, | |
| "quarter": q, | |
| "mode": mode, | |
| **m, | |
| } | |
| ) | |
| feat_sets[(L, q, mode)] = feats | |
| print(f"[{axis}] finished layer {L}", flush=True) | |
| # aggregate | |
| from collections import defaultdict | |
| grouped = defaultdict(list) | |
| for r in rows: | |
| grouped[(r["layer"], r["quarter"], r["mode"])].append(r) | |
| METRIC_KEYS = [ | |
| "prediction_mse", | |
| "prediction_mse_ratio", | |
| "target_cosine", | |
| "prediction_shift_l2", | |
| "clean_prediction_cosine", | |
| "substitution_maha", | |
| "prediction_maha", | |
| ] | |
| summaries = [] | |
| for (L, q, mode), g in sorted(grouped.items()): | |
| s = {"layer": L, "quarter": q, "mode": mode, "n": len(g)} | |
| for mk in METRIC_KEYS: | |
| vals = np.array([r[mk] for r in g if r.get(mk) is not None], float) | |
| s["mean_" + mk] = float(vals.mean()) if vals.size else None | |
| s["std_" + mk] = float(vals.std(ddof=1)) if vals.size > 1 else 0.0 | |
| summaries.append(s) | |
| # probes | |
| prober = LatentProber(seed=SEED) | |
| probe_cache = {} | |
| clean_acc = None | |
| total_probes = len(feat_sets) | |
| trained = 0 | |
| for (L, q, mode), feats in feat_sets.items(): | |
| if mode == "clean": | |
| if clean_acc is None: | |
| r = prober.train_probe( | |
| activations=torch.tensor(feats, dtype=torch.float32), | |
| labels=np.array(labels, np.int64), | |
| concept_name="focus_class", | |
| activation_name="clean", | |
| probe_type="logistic", | |
| test_split=PROBE_SPLIT, | |
| use_cv=False, | |
| ) | |
| clean_acc = float(r.accuracy) | |
| probe_cache[(L, q, mode)] = clean_acc | |
| else: | |
| r = prober.train_probe( | |
| activations=torch.tensor(feats, dtype=torch.float32), | |
| labels=np.array(labels, np.int64), | |
| concept_name="focus_class", | |
| activation_name=f"{axis}_L{L}_q{q}_{mode}", | |
| probe_type="logistic", | |
| test_split=PROBE_SPLIT, | |
| use_cv=False, | |
| ) | |
| probe_cache[(L, q, mode)] = float(r.accuracy) | |
| trained += 1 | |
| if trained % 20 == 0 or trained == total_probes: | |
| print(f"[{axis}] probe {trained}/{total_probes} trained", flush=True) | |
| for s in summaries: | |
| s["classification_accuracy"] = probe_cache[(s["layer"], s["quarter"], s["mode"])] | |
| # write data | |
| cfg = { | |
| "axis": axis, | |
| "n_samples": N, | |
| "classes": FOLDERS, | |
| "human": HUMAN, | |
| "layers": layers, | |
| "n_quarters": NQ, | |
| "quarter_meta": meta, | |
| "modes": MODES, | |
| "batch": BATCH, | |
| "seed": SEED, | |
| "precision": str(DTYPE), | |
| "run_id": ts, | |
| "clean_probe_accuracy": clean_acc, | |
| } | |
| (run_dir / "config.yaml").write_text(yaml.safe_dump(cfg, sort_keys=False)) | |
| (run_dir / "summary_metrics.json").write_text(json.dumps(summaries, indent=2)) | |
| (run_dir / "per_sample_metrics.json").write_text(json.dumps(rows, indent=2)) | |
| make_plots(run_dir, axis, summaries, layers, clean_acc) | |
| print(f"[{axis}] saved -> {run_dir}", flush=True) | |
| return run_dir | |
| def make_plots(run_dir, axis, summaries, layers, clean_acc): | |
| look = {(s["layer"], s["quarter"], s["mode"]): s for s in summaries} | |
| quarters = list(range(NQ)) | |
| plots = [ | |
| ("mean_prediction_mse", "Prediction MSE", (0, None)), | |
| ("mean_target_cosine", "Target cosine (0-1)", (0, 1)), | |
| ("classification_accuracy", "Linear-probe accuracy (0-1)", (0, 1)), | |
| ("mean_substitution_maha", "Substitution Mahalanobis", (0, None)), | |
| ("mean_prediction_maha", "Prediction Mahalanobis", (0, None)), | |
| ] | |
| for metric, ylabel, ylim in plots: | |
| fig, axes = plt.subplots(1, 3, figsize=(20, 5.5), sharey=True) | |
| for mi, mode in enumerate(MODES): | |
| ax = axes[mi] | |
| width = 0.8 / NQ | |
| x = np.arange(len(layers)) | |
| for q in quarters: | |
| vals = [look.get((L, q, mode), {}).get(metric) for L in layers] | |
| vals = [np.nan if v is None else v for v in vals] | |
| ax.bar(x - 0.4 + width / 2 + q * width, vals, width=width, label=f"Q{q}") | |
| if metric in ("mean_prediction_mse", "mean_target_cosine", "classification_accuracy"): | |
| base = ( | |
| clean_acc | |
| if metric == "classification_accuracy" | |
| else look.get((layers[0], -1, "clean"), {}).get(metric) | |
| ) | |
| if base is not None: | |
| ax.axhline( | |
| base, | |
| ls="--", | |
| lw=1, | |
| color="k", | |
| alpha=0.6, | |
| label="clean" if mi == 0 else None, | |
| ) | |
| ax.set_title(f"{axis} · {mode}") | |
| ax.set_xticks(x, [str(L) for L in layers]) | |
| ax.set_xlabel("Predictor layer") | |
| if ylim: | |
| ax.set_ylim(*ylim) | |
| ax.grid(axis="y", alpha=0.25) | |
| if mi == 0: | |
| ax.set_ylabel(ylabel) | |
| ax.legend(frameon=False, ncol=2, fontsize=8) | |
| fig.suptitle(f"I-JEPA quarter-ablation ({axis} axis): {ylabel}") | |
| fig.tight_layout() | |
| fig.savefig(run_dir / f"{metric}.png", dpi=170, bbox_inches="tight") | |
| plt.close(fig) | |
| def main(): | |
| seed_all() | |
| a = load_model() | |
| context_ids, target_ids, n_patches = fixed_masks(a) | |
| layers = list(range(len(a.predictor.blocks))) | |
| print(f"model ready: patches={n_patches} depth={len(layers)}", flush=True) | |
| clean = collect_clean(a, context_ids, target_ids, layers) | |
| n_tokens = clean["act"][0].shape[1] | |
| n_hidden = clean["act"][0].shape[2] | |
| print(f"clean cached: act shape per layer = {tuple(clean['act'][0].shape)}", flush=True) | |
| layer_means, ldet = {}, {} | |
| for L in layers: | |
| act = clean["act"][L].float() | |
| layer_means[L] = act.mean(0) # [seq, hidden] | |
| ldet[L] = MahalanobisOODDetector().fit(act.mean(1)) # pooled over tokens -> [hidden] | |
| pdet = MahalanobisOODDetector().fit(clean["pred"].float().mean(1)) | |
| out = {} | |
| for axis in ["token", "hidden"]: | |
| out[axis] = str( | |
| run_axis( | |
| a, | |
| axis, | |
| context_ids, | |
| target_ids, | |
| layers, | |
| clean, | |
| layer_means, | |
| ldet, | |
| pdet, | |
| n_tokens, | |
| n_hidden, | |
| ) | |
| ) | |
| print("ALL DONE", json.dumps(out), flush=True) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 17.6 kB
- Xet hash:
- c1b1fa9540b569a649efc02f05ec37b3bce4471c5e14284b5e0cc903b3f6f575
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.