Rishik001's picture
download
raw
17.6 kB
#!/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
@torch.no_grad()
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},
}
@torch.no_grad()
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.