Buckets:
| #!/usr/bin/env python3 | |
| """Task 4 deep dive. | |
| RUN A (exp 3, 180 focus samples): | |
| * clean analyses: per-layer anisotropy, attention-mass received by the 16 mask | |
| tokens, and a prediction-lens (snap the readout head onto every layer). | |
| * Q3 fine split: ablate the 47 context tokens (192-239) vs the 16 target/mask | |
| tokens (240-255), all layers, all 3 modes, with probe + shrinkage Mahalanobis. | |
| * per-target-token: ablate each of the 16 target tokens on its own, all layers, | |
| all modes; aggregate per class (which target token matters for which class). | |
| RUN B (exp 1, full 1000 samples / 50 classes): | |
| * token quarters + hidden quarters + Q3 split, all layers, all 3 modes, probe | |
| (no CV for speed) + shrinkage Mahalanobis. | |
| Mahalanobis uses Ledoit-Wolf shrinkage so it is well defined when the feature | |
| dimension exceeds the sample count (fixes the earlier nan on 1280-dim predictions). | |
| """ | |
| # ruff: noqa: N802, N803, N806 -- names mirror this analysis's linear-algebra notation | |
| # (e.g. X, L, N, B, M for matrices/batch/layer/count) rather than PEP8 casing. | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import time | |
| from collections import defaultdict | |
| 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 | |
| from sklearn.covariance import LedoitWolf | |
| 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.data import load_imagenet_image # noqa: E402 | |
| from world_model_lens.hub.model_hub import ModelHub # noqa: E402 | |
| SCRATCH = Path(__file__).resolve().parent | |
| OUT = Path(os.environ.get("WML_OUTPUT_ROOT", str(TASK_ROOT / "outputs"))) | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32 | |
| BATCH = 32 | |
| SEED = 42 | |
| MODES = ["zero", "mean", "resample"] | |
| TARGET_SIDE = 4 | |
| 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 masks(a): | |
| n = int(a.context_encoder.patch_embed.n_patches) | |
| grid = int(math.sqrt(n)) | |
| side = TARGET_SIDE | |
| s = (grid - side) // 2 | |
| target = {r * grid + c for r in range(s, s + side) for c in range(s, s + side)} | |
| context = [p for p in range(n) if p not in target] | |
| return context, sorted(target), n | |
| class ShrinkMaha: | |
| """Mahalanobis with Ledoit-Wolf shrinkage precision (works when D > N).""" | |
| def fit(self, X): | |
| X = np.asarray(X, np.float64) | |
| self.mu = X.mean(0) | |
| self.prec = LedoitWolf().fit(X).precision_ | |
| return self | |
| def score(self, x): | |
| d = np.asarray(x, np.float64).reshape(-1) - self.mu | |
| return float(np.sqrt(max(float(d @ self.prec @ d), 0.0))) | |
| def collect_clean(a, samples, context_ids, target_ids, layers): | |
| store = {} | |
| def cap(L): | |
| def h(_m, _i, o): | |
| store[L] = o.detach() | |
| return h | |
| ctx, pred, tgt = [], [], [] | |
| acts = {L: [] for L in layers} | |
| for s in range(0, len(samples), BATCH): | |
| idx = list(range(s, min(s + BATCH, len(samples)))) | |
| obs = torch.cat( | |
| [load_imagenet_image(samples[i]["path"], image_size=224) for i in idx], 0 | |
| ).to(DEVICE, DTYPE) | |
| hs = [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) | |
| pr = a.predictor(cl, context_ids, target_ids) | |
| tg = a.target_encoder(obs) | |
| finally: | |
| for h in hs: | |
| h.remove() | |
| ctx.append(cl.to("cpu", torch.float16)) | |
| pred.append(pr.to("cpu", torch.float16)) | |
| tgt.append(tg[:, target_ids, :].to("cpu", torch.float16)) | |
| for L in layers: | |
| acts[L].append(store[L].to("cpu", torch.float16)) | |
| return { | |
| "context": torch.cat(ctx), | |
| "pred": torch.cat(pred), | |
| "tgt": torch.cat(tgt), | |
| "acts": {L: torch.cat(acts[L]) for L in layers}, | |
| } | |
| def clean_analyses(a, clean, layers, n_context, n_target): | |
| """Anisotropy, attention-to-mask, and prediction-lens on clean activations.""" | |
| res = {"layers": layers} | |
| # prediction-lens: snap norm+project_back onto each layer's mask tokens | |
| norm = a.predictor.norm | |
| proj = a.predictor.predictor_project_back | |
| tgt = clean["tgt"].float() # [N,16,1280] | |
| lens_cos = {} | |
| for L in layers: | |
| mtok = clean["acts"][L][:, n_context:, :].to(DEVICE, DTYPE) # [N,16,384] | |
| read = proj(norm(mtok)).float().cpu() # [N,16,1280] | |
| c = torch.nn.functional.cosine_similarity(read.flatten(1), tgt.flatten(1), dim=1) | |
| lens_cos[L] = float(c.mean()) | |
| res["prediction_lens_cos_to_target"] = lens_cos | |
| # anisotropy: eigenspectrum of token activations per layer | |
| aniso = {} | |
| for L in layers: | |
| X = clean["acts"][L].float().reshape(-1, clean["acts"][L].shape[-1]) # [N*256,384] | |
| if X.shape[0] > 6000: | |
| sel = torch.randperm(X.shape[0])[:6000] | |
| X = X[sel] | |
| Xc = X - X.mean(0) | |
| cov = (Xc.T @ Xc) / (Xc.shape[0] - 1) | |
| ev = torch.linalg.eigvalsh(cov).clamp_min(0).flip(0).numpy() | |
| pr = float((ev.sum() ** 2) / (np.sum(ev**2) + 1e-12)) # participation ratio | |
| p = ev / (ev.sum() + 1e-12) | |
| eff_rank = float(np.exp(-(p * np.log(p + 1e-12)).sum())) # effective rank | |
| aniso[L] = { | |
| "participation_ratio": pr, | |
| "effective_rank": eff_rank, | |
| "top1_var_frac": float(ev[0] / (ev.sum() + 1e-12)), | |
| "dim": int(ev.shape[0]), | |
| } | |
| res["anisotropy"] = aniso | |
| # attention mass received by the 16 mask tokens (cols n_context:), per layer | |
| n_seq = n_context + n_target | |
| n_attn = min(8, len(clean["_samples"])) | |
| obs = torch.cat( | |
| [load_imagenet_image(clean["_samples"][i]["path"], image_size=224) for i in range(n_attn)], | |
| 0, | |
| ).to(DEVICE, DTYPE) | |
| _ = a.predictor(a.context_encoder(obs, patch_ids=clean["_ctx"]), clean["_ctx"], clean["_tgt"]) | |
| attn_recv = {} | |
| attn_from_mask = {} | |
| for L in layers: | |
| w = a.predictor.blocks[L].attn.last_attn_weights.float() # [B,H,seq,seq] | |
| recv = w[:, :, :, n_context:].sum(-1).mean().item() # mass any query sends to mask cols | |
| frm = ( | |
| w[:, :, n_context:, :n_context].sum(-1).mean().item() | |
| ) # mass mask rows send to context | |
| attn_recv[L] = recv | |
| attn_from_mask[L] = frm | |
| res["attn_mass_to_mask_tokens"] = attn_recv | |
| res["attn_mass_mask_to_context"] = attn_from_mask | |
| res["uniform_baseline_to_mask"] = n_target / n_seq | |
| return res | |
| def intervene(a, ctx_batch, context_ids, target_ids, L, hook): | |
| h = a.predictor.blocks[L].hook_resid_post.register_forward_hook(hook) | |
| try: | |
| pr = a.predictor(ctx_batch, context_ids, target_ids) | |
| finally: | |
| h.remove() | |
| return pr.detach() | |
| def make_token_hook(tok_idx, mode, mean_dev, donor_dev): | |
| ti = torch.as_tensor(tok_idx, dtype=torch.long) | |
| def hook(_m, _i, out): | |
| o = out.clone() | |
| if mode == "zero": | |
| o[:, ti, :] = 0 | |
| elif mode == "mean": | |
| o[:, ti, :] = mean_dev[:, ti, :].to(o.dtype) | |
| else: | |
| o[:, ti, :] = donor_dev[:, ti, :].to(o.dtype) | |
| return o | |
| return hook | |
| def make_hidden_hook(sl, mode, mean_dev, donor_dev): | |
| def hook(_m, _i, out): | |
| o = out.clone() | |
| 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, ldet, pdet): | |
| 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": ldet.score(mod_feat), | |
| "prediction_maha": pdet.score(pred.float().mean(0)), | |
| } | |
| def run_conditions( | |
| a, | |
| samples, | |
| context_ids, | |
| target_ids, | |
| layers, | |
| clean, | |
| layer_means, | |
| ldet, | |
| pdet, | |
| conditions, | |
| labels, | |
| collect_feats, | |
| N, | |
| ): | |
| """conditions: list of dicts {name, kind: 'token'|'hidden', sel, mode}. | |
| sel = token-index list (kind=token) or slice (kind=hidden). | |
| Returns (rows, feat_map) where feat_map[(name,L,mode)] = pooled features if | |
| collect_feats else {}.""" | |
| rows = [] | |
| feat_map = {} | |
| batches = [list(range(s, min(s + BATCH, N))) for s in range(0, N, BATCH)] | |
| for cond in conditions: | |
| name, kind, sel, mode = cond["name"], cond["kind"], cond["sel"], cond["mode"] | |
| for L in layers: | |
| feats = [] | |
| lm = layer_means[L] | |
| ld = ldet[L] | |
| ca = clean["acts"][L].float() | |
| 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) | |
| donor_idx = [(i + 1) % N for i in batch] if mode == "resample" else None | |
| donor_dev = ca[donor_idx].to(DEVICE) if donor_idx else None | |
| if kind == "token": | |
| hook = make_token_hook(sel, mode, mean_dev, donor_dev) | |
| else: | |
| hook = make_hidden_hook(sel, 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()) | |
| mod = ca[batch].clone() | |
| if kind == "token": | |
| ti = torch.as_tensor(sel, dtype=torch.long) | |
| if mode == "zero": | |
| mod[:, ti, :] = 0 | |
| elif mode == "mean": | |
| mod[:, ti, :] = lm[ti, :].unsqueeze(0) | |
| else: | |
| mod[:, ti, :] = ca[donor_idx][:, ti, :] | |
| else: | |
| if mode == "zero": | |
| mod[:, :, sel] = 0 | |
| elif mode == "mean": | |
| mod[:, :, sel] = lm[:, sel].unsqueeze(0) | |
| else: | |
| mod[:, :, sel] = ca[donor_idx][:, :, sel] | |
| mfeat = mod.mean(1) | |
| for pos, i in enumerate(batch): | |
| m = metrics( | |
| pred[pos], | |
| clean["tgt"][i].float(), | |
| clean["pred"][i].float(), | |
| mfeat[pos], | |
| ld, | |
| pdet, | |
| ) | |
| rows.append( | |
| { | |
| "sample_index": i, | |
| "label": labels[i], | |
| "class_name": samples[i]["class_name"], | |
| "human": HUMAN.get(samples[i]["class_name"], samples[i]["class_name"]), | |
| "condition": name, | |
| "layer": L, | |
| "mode": mode, | |
| **m, | |
| } | |
| ) | |
| if collect_feats: | |
| feat_map[(name, L, mode)] = feats | |
| print(f" cond {name}/{mode} done", flush=True) | |
| return rows, feat_map | |
| def attach_probes(summ, feat_map, labels, use_cv): | |
| for s in summ: | |
| key = (s["condition"], s["layer"], s["mode"]) | |
| if key in feat_map: | |
| s["classification_accuracy"] = probe_for( | |
| feat_map[key], labels, f"{key[0]}_L{key[1]}_{key[2]}", use_cv | |
| ) | |
| def aggregate(rows, extra_keys=("condition",)): | |
| grouped = defaultdict(list) | |
| for r in rows: | |
| grouped[(r["condition"], r["layer"], r["mode"])].append(r) | |
| MK = [ | |
| "prediction_mse", | |
| "prediction_mse_ratio", | |
| "target_cosine", | |
| "prediction_shift_l2", | |
| "clean_prediction_cosine", | |
| "substitution_maha", | |
| "prediction_maha", | |
| ] | |
| out = [] | |
| for (cond, L, mode), g in sorted(grouped.items()): | |
| s = {"condition": cond, "layer": L, "mode": mode, "n": len(g)} | |
| for mk in MK: | |
| v = np.array([r[mk] for r in g if r.get(mk) is not None], float) | |
| s["mean_" + mk] = float(v.mean()) if v.size else None | |
| s["std_" + mk] = float(v.std(ddof=1)) if v.size > 1 else 0.0 | |
| out.append(s) | |
| return out | |
| def probe_for(a_feats, labels, name, use_cv): | |
| r = LatentProber(seed=SEED).train_probe( | |
| activations=torch.tensor(a_feats, dtype=torch.float32), | |
| labels=np.asarray(labels, np.int64), | |
| concept_name="cls", | |
| activation_name=name, | |
| probe_type="logistic", | |
| test_split=0.2, | |
| use_cv=use_cv, | |
| ) | |
| return float(r.accuracy) | |
| def load_focus_180(): | |
| return json.load(open(SCRATCH / "focus_180_official.json")) | |
| def load_full_1000(): | |
| man = json.load(open(SCRATCH / "focus_1000_official.json")) | |
| return man | |
| # --------------------------------------------------------------------------- RUN A | |
| def run_A(a, context_ids, target_ids, layers, n_ctx, n_tgt): | |
| ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") | |
| rd = OUT / "ijepa_task4_deep" / "expA_q3split_180" / ts | |
| rd.mkdir(parents=True, exist_ok=True) | |
| samples = load_focus_180() | |
| N = len(samples) | |
| folders = sorted({s["class_name"] for s in samples}, key=lambda x: int(x)) | |
| f2l = {f: i for i, f in enumerate(folders)} | |
| labels = [f2l[s["class_name"]] for s in samples] | |
| print(f"[A] {N} samples, {len(folders)} classes", flush=True) | |
| clean = collect_clean(a, samples, context_ids, target_ids, layers) | |
| clean["_samples"] = samples | |
| clean["_ctx"] = context_ids | |
| clean["_tgt"] = target_ids | |
| layer_means = {L: clean["acts"][L].float().mean(0) for L in layers} | |
| ldet = {L: ShrinkMaha().fit(clean["acts"][L].float().mean(1).numpy()) for L in layers} | |
| pdet = ShrinkMaha().fit(clean["pred"].float().mean(1).numpy()) | |
| ca = clean_analyses(a, clean, layers, n_ctx, n_tgt) | |
| (rd / "clean_analyses.json").write_text(json.dumps(ca, indent=2)) | |
| print("[A] clean analyses done", flush=True) | |
| ctx_q3 = list(range(192, n_ctx)) # 192-239 context tokens inside Q3 | |
| tgt_tokens = list(range(n_ctx, n_ctx + n_tgt)) # 240-255 mask tokens | |
| conds = [] | |
| for mode in MODES: | |
| conds.append({"name": "Q3_context47", "kind": "token", "sel": ctx_q3, "mode": mode}) | |
| conds.append({"name": "Q3_target16", "kind": "token", "sel": tgt_tokens, "mode": mode}) | |
| rows, fmap = run_conditions( | |
| a, | |
| samples, | |
| context_ids, | |
| target_ids, | |
| layers, | |
| clean, | |
| layer_means, | |
| ldet, | |
| pdet, | |
| conds, | |
| labels, | |
| collect_feats=True, | |
| N=N, | |
| ) | |
| print("[A] Q3 split done", flush=True) | |
| # per-target-token: each single mask token | |
| conds_tok = [] | |
| for mode in MODES: | |
| for k, tok in enumerate(tgt_tokens): | |
| conds_tok.append({"name": f"tgt_tok_{k}", "kind": "token", "sel": [tok], "mode": mode}) | |
| rows_tok, _ = run_conditions( | |
| a, | |
| samples, | |
| context_ids, | |
| target_ids, | |
| layers, | |
| clean, | |
| layer_means, | |
| ldet, | |
| pdet, | |
| conds_tok, | |
| labels, | |
| collect_feats=False, | |
| N=N, | |
| ) | |
| print("[A] per-target-token done", flush=True) | |
| summ = aggregate(rows) | |
| attach_probes(summ, fmap, labels, use_cv=True) | |
| summ_tok = aggregate(rows_tok) | |
| (rd / "summary_q3split.json").write_text(json.dumps(summ, indent=2)) | |
| (rd / "summary_target_tokens.json").write_text(json.dumps(summ_tok, indent=2)) | |
| (rd / "per_sample_target_tokens.json").write_text(json.dumps(rows_tok, indent=2)) | |
| (rd / "per_sample_q3split.json").write_text(json.dumps(rows, indent=2)) | |
| plots_A(rd, summ, summ_tok, ca, layers, samples, rows_tok, labels) | |
| (rd / "config.yaml").write_text( | |
| yaml.safe_dump( | |
| { | |
| "run": "A", | |
| "n": N, | |
| "classes": folders, | |
| "layers": layers, | |
| "n_context": n_ctx, | |
| "n_target": n_tgt, | |
| "modes": MODES, | |
| "seed": SEED, | |
| "maha": "ledoit-wolf shrinkage", | |
| }, | |
| sort_keys=False, | |
| ) | |
| ) | |
| print(f"[A] saved -> {rd}", flush=True) | |
| return str(rd) | |
| def plots_A(rd, summ, summ_tok, ca, layers, samples, rows_tok, labels): | |
| look = {(s["condition"], s["layer"], s["mode"]): s for s in summ} | |
| # Q3 split: context47 vs target16, MSE and cosine, per mode | |
| for metric, ylab, ylim in [ | |
| ("mean_prediction_mse", "Prediction MSE", (0, None)), | |
| ("mean_target_cosine", "Target cosine", (0, 1)), | |
| ("mean_substitution_maha", "Substitution Maha (shrinkage)", (0, None)), | |
| ("mean_prediction_maha", "Prediction Maha (shrinkage)", (0, None)), | |
| ]: | |
| fig, axes = plt.subplots(1, 3, figsize=(18, 5), sharey=True) | |
| for mi, mode in enumerate(MODES): | |
| ax = axes[mi] | |
| x = np.arange(len(layers)) | |
| w = 0.38 | |
| for j, cond in enumerate(["Q3_context47", "Q3_target16"]): | |
| vals = [look.get((cond, L, mode), {}).get(metric) for L in layers] | |
| vals = [np.nan if v is None else v for v in vals] | |
| ax.bar(x - w / 2 + j * w, vals, width=w, label=cond) | |
| ax.set_title(mode) | |
| ax.set_xticks(x, [str(L) for L in layers]) | |
| ax.set_xlabel("layer") | |
| if ylim: | |
| ax.set_ylim(*ylim) | |
| ax.grid(axis="y", alpha=0.25) | |
| if mi == 0: | |
| ax.set_ylabel(ylab) | |
| ax.legend(frameon=False, fontsize=8) | |
| fig.suptitle(f"Q3 split (47 context vs 16 target tokens): {ylab}") | |
| fig.tight_layout() | |
| fig.savefig(rd / f"q3split_{metric}.png", dpi=160) | |
| plt.close(fig) | |
| # anisotropy + attention + lens (clean) | |
| fig, axes = plt.subplots(1, 3, figsize=(18, 4.5)) | |
| L = ca["layers"] | |
| axes[0].plot( | |
| L, | |
| [ | |
| ca["anisotropy"][ | |
| str(layer) if isinstance(list(ca["anisotropy"].keys())[0], str) else layer | |
| ]["effective_rank"] | |
| for layer in L | |
| ], | |
| "o-", | |
| ) | |
| axes[0].set_title("Effective rank (higher = more isotropic)") | |
| axes[0].set_xlabel("layer") | |
| axes[0].grid(alpha=0.3) | |
| axes[1].plot( | |
| L, | |
| [ | |
| ca["attn_mass_to_mask_tokens"][ | |
| ( | |
| str(layer) | |
| if isinstance(list(ca["attn_mass_to_mask_tokens"].keys())[0], str) | |
| else layer | |
| ) | |
| ] | |
| for layer in L | |
| ], | |
| "o-", | |
| label="to mask", | |
| ) | |
| axes[1].axhline(ca["uniform_baseline_to_mask"], ls="--", color="k", alpha=0.6, label="uniform") | |
| axes[1].set_title("Attention mass received by 16 mask tokens") | |
| axes[1].set_xlabel("layer") | |
| axes[1].legend() | |
| axes[1].grid(alpha=0.3) | |
| axes[2].plot( | |
| L, | |
| [ | |
| ca["prediction_lens_cos_to_target"][ | |
| ( | |
| str(layer) | |
| if isinstance(list(ca["prediction_lens_cos_to_target"].keys())[0], str) | |
| else layer | |
| ) | |
| ] | |
| for layer in L | |
| ], | |
| "o-", | |
| ) | |
| axes[2].set_title("Prediction-lens cosine to target") | |
| axes[2].set_xlabel("layer") | |
| axes[2].set_ylim(0, 1) | |
| axes[2].grid(alpha=0.3) | |
| fig.tight_layout() | |
| fig.savefig(rd / "clean_analyses.png", dpi=160) | |
| plt.close(fig) | |
| # per-target-token heatmap: token (16) x class (mean MSE, zero mode, worst layer) | |
| # aggregate rows_tok by (class, token) for zero mode averaged over layers | |
| by = defaultdict(list) | |
| for r in rows_tok: | |
| if r["mode"] == "zero": | |
| by[(r["human"], r["condition"])].append(r["prediction_mse"]) | |
| classes = sorted({r["human"] for r in rows_tok}) | |
| toks = [f"tgt_tok_{k}" for k in range(16)] | |
| M = np.array([[np.mean(by[(c, t)]) if by[(c, t)] else np.nan for t in toks] for c in classes]) | |
| fig, ax = plt.subplots(figsize=(11, 6)) | |
| im = ax.imshow(M, aspect="auto", cmap="viridis") | |
| ax.set_xticks(range(16), [str(k) for k in range(16)]) | |
| ax.set_yticks(range(len(classes)), classes) | |
| ax.set_xlabel("target token index (0-15)") | |
| ax.set_title("Per-target-token zero: prediction MSE by class (avg over layers)") | |
| fig.colorbar(im, ax=ax) | |
| fig.tight_layout() | |
| fig.savefig(rd / "target_token_by_class.png", dpi=160) | |
| plt.close(fig) | |
| # --------------------------------------------------------------------------- RUN B | |
| def run_B(a, context_ids, target_ids, layers, n_ctx, n_tgt, n_seq, n_hidden): | |
| ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") | |
| rd = OUT / "ijepa_task4_deep" / "expB_full1000" / ts | |
| rd.mkdir(parents=True, exist_ok=True) | |
| samples = load_full_1000() | |
| N = len(samples) | |
| folders = sorted({s["class_name"] for s in samples}) | |
| f2l = {f: i for i, f in enumerate(folders)} | |
| labels = [f2l[s["class_name"]] for s in samples] | |
| print(f"[B] {N} samples, {len(folders)} classes", flush=True) | |
| clean = collect_clean(a, samples, context_ids, target_ids, layers) | |
| clean["_samples"] = samples | |
| clean["_ctx"] = context_ids | |
| clean["_tgt"] = target_ids | |
| layer_means = {L: clean["acts"][L].float().mean(0) for L in layers} | |
| ldet = {L: ShrinkMaha().fit(clean["acts"][L].float().mean(1).numpy()) for L in layers} | |
| pdet = ShrinkMaha().fit(clean["pred"].float().mean(1).numpy()) | |
| ca = clean_analyses(a, clean, layers, n_ctx, n_tgt) | |
| (rd / "clean_analyses.json").write_text(json.dumps(ca, indent=2)) | |
| print("[B] clean analyses done", flush=True) | |
| tq = n_seq // 4 | |
| hq = n_hidden // 4 | |
| conds = [] | |
| for mode in MODES: | |
| for q in range(4): | |
| conds.append( | |
| { | |
| "name": f"token_Q{q}", | |
| "kind": "token", | |
| "sel": list(range(q * tq, (q + 1) * tq)), | |
| "mode": mode, | |
| } | |
| ) | |
| for q in range(4): | |
| conds.append( | |
| { | |
| "name": f"hidden_Q{q}", | |
| "kind": "hidden", | |
| "sel": slice(q * hq, (q + 1) * hq), | |
| "mode": mode, | |
| } | |
| ) | |
| conds.append( | |
| {"name": "Q3_context47", "kind": "token", "sel": list(range(192, n_ctx)), "mode": mode} | |
| ) | |
| conds.append( | |
| { | |
| "name": "Q3_target16", | |
| "kind": "token", | |
| "sel": list(range(n_ctx, n_ctx + n_tgt)), | |
| "mode": mode, | |
| } | |
| ) | |
| rows, fmap = run_conditions( | |
| a, | |
| samples, | |
| context_ids, | |
| target_ids, | |
| layers, | |
| clean, | |
| layer_means, | |
| ldet, | |
| pdet, | |
| conds, | |
| labels, | |
| collect_feats=True, | |
| N=N, | |
| ) | |
| print("[B] conditions done", flush=True) | |
| summ = aggregate(rows) | |
| attach_probes(summ, fmap, labels, use_cv=False) | |
| (rd / "summary_metrics.json").write_text(json.dumps(summ, indent=2)) | |
| (rd / "per_sample_metrics.json").write_text(json.dumps(rows, indent=2)) | |
| (rd / "config.yaml").write_text( | |
| yaml.safe_dump( | |
| { | |
| "run": "B", | |
| "n": N, | |
| "classes": folders, | |
| "layers": layers, | |
| "modes": MODES, | |
| "seed": SEED, | |
| "maha": "ledoit-wolf shrinkage", | |
| "probe": "skipped-for-speed", | |
| }, | |
| sort_keys=False, | |
| ) | |
| ) | |
| plots_B(rd, summ, ca, layers) | |
| print(f"[B] saved -> {rd}", flush=True) | |
| return str(rd) | |
| def plots_B(rd, summ, ca, layers): | |
| look = {(s["condition"], s["layer"], s["mode"]): s for s in summ} | |
| for metric, ylab, ylim in [ | |
| ("mean_prediction_mse", "Prediction MSE", (0, None)), | |
| ("mean_target_cosine", "Target cosine", (0, 1)), | |
| ("mean_substitution_maha", "Substitution Maha (shrinkage)", (0, None)), | |
| ("mean_prediction_maha", "Prediction Maha (shrinkage)", (0, None)), | |
| ]: | |
| for group, conds in [ | |
| ("token", [f"token_Q{q}" for q in range(4)]), | |
| ("hidden", [f"hidden_Q{q}" for q in range(4)]), | |
| ]: | |
| fig, axes = plt.subplots(1, 3, figsize=(18, 5), sharey=True) | |
| for mi, mode in enumerate(MODES): | |
| ax = axes[mi] | |
| x = np.arange(len(layers)) | |
| w = 0.8 / 4 | |
| for j, cond in enumerate(conds): | |
| vals = [look.get((cond, L, 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 + w / 2 + j * w, vals, width=w, label=f"Q{j}") | |
| ax.set_title(mode) | |
| ax.set_xticks(x, [str(L) for L in layers]) | |
| ax.set_xlabel("layer") | |
| if ylim: | |
| ax.set_ylim(*ylim) | |
| ax.grid(axis="y", alpha=0.25) | |
| if mi == 0: | |
| ax.set_ylabel(ylab) | |
| ax.legend(frameon=False, fontsize=8, ncol=2) | |
| fig.suptitle(f"[full 1000] {group} quarters: {ylab}") | |
| fig.tight_layout() | |
| fig.savefig(rd / f"{group}_{metric}.png", dpi=160) | |
| plt.close(fig) | |
| def main(): | |
| seed_all() | |
| a = load_model() | |
| context_ids, target_ids, n_patches = masks(a) | |
| layers = list(range(len(a.predictor.blocks))) | |
| n_ctx = len(context_ids) | |
| n_tgt = len(target_ids) | |
| n_seq = n_patches | |
| n_hidden = a.predictor.predictor_embed_dim | |
| print( | |
| f"model ready: seq={n_seq} ctx={n_ctx} tgt={n_tgt} hidden={n_hidden} depth={len(layers)}", | |
| flush=True, | |
| ) | |
| out = {} | |
| t0 = time.time() | |
| out["A"] = run_A(a, context_ids, target_ids, layers, n_ctx, n_tgt) | |
| print(f"=== RUN A complete in {time.time()-t0:.0f}s ===", flush=True) | |
| t1 = time.time() | |
| out["B"] = run_B(a, context_ids, target_ids, layers, n_ctx, n_tgt, n_seq, n_hidden) | |
| print(f"=== RUN B complete in {time.time()-t1:.0f}s ===", flush=True) | |
| print("ALL DONE", json.dumps(out), flush=True) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 27.3 kB
- Xet hash:
- 65752d15c735d083d156bc28db2fa25ace3bb864292cad2a9bcdcbec0d7b305a
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.