Rishik001/ijepa_stuff / task47_0809 /scripts /run_resample_class.py
Rishik001's picture
download
raw
10.6 kB
#!/usr/bin/env python3
"""Resample ablation controlled by donor class similarity.
Donor for the injected activation is drawn from one of three pools, per sample:
* same -- another image of the SAME class
* similar -- an image of the nearest OTHER class (by clean-embedding centroid)
* different -- an image of the farthest class
Class similarity is computed data-drivenly from clean target-encoder embeddings.
Two intervention loci on the token axis:
* target16 -- replace only the 16 mask tokens' activation (the meaningful locus)
* full -- replace the whole residual state (layer-independent: donor's pred)
180 focus samples. Metrics: MSE, target cosine, probe accuracy, shrinkage
Mahalanobis. Hypothesis: same-class donor is gentlest, different-class harshest.
"""
# ruff: noqa: N806 -- names mirror this analysis's linear-algebra notation
# (e.g. N for sample count, C for class count) rather than PEP8 casing.
from __future__ import annotations
import json
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
SP = Path(__file__).resolve().parent
sys.path.insert(0, str(SP))
from run_deep import ( # noqa: E402
BATCH,
DEVICE,
DTYPE,
HUMAN,
OUT,
SEED,
ShrinkMaha,
aggregate,
collect_clean,
intervene,
load_focus_180,
load_model,
masks,
metrics,
probe_for,
seed_all,
)
LEVELS = ["same", "similar", "different"]
def build_donor_maps(samples, clean, folders):
"""Return donor_map[level] -> np.int array [N] of donor sample indices, plus the
per-class nearest/farthest class picks."""
N = len(samples)
labels = np.array([folders.index(s["class_name"]) for s in samples])
emb = clean["tgt"].float().mean(1).numpy() # [N,1280] pooled target emb
emb = emb / (np.linalg.norm(emb, axis=1, keepdims=True) + 1e-8)
C = len(folders)
cent = np.stack([emb[labels == c].mean(0) for c in range(C)])
cent = cent / (np.linalg.norm(cent, axis=1, keepdims=True) + 1e-8)
sim = cent @ cent.T # [C,C] cosine
sim_near = sim.copy()
np.fill_diagonal(sim_near, -np.inf)
sim_far = sim.copy()
np.fill_diagonal(sim_far, np.inf)
nearest = sim_near.argmax(1)
farthest = sim_far.argmin(1)
idx_by_class = {c: np.where(labels == c)[0] for c in range(C)}
rng = np.random.default_rng(SEED)
donor = {lv: np.zeros(N, dtype=int) for lv in LEVELS}
for i in range(N):
c = labels[i]
same_pool = idx_by_class[c][idx_by_class[c] != i]
donor["same"][i] = rng.choice(same_pool)
donor["similar"][i] = rng.choice(idx_by_class[nearest[c]])
donor["different"][i] = rng.choice(idx_by_class[farthest[c]])
picks = {
folders[c]: {
"similar_class": folders[nearest[c]],
"different_class": folders[farthest[c]],
"sim_to_similar": float(sim[c, nearest[c]]),
"sim_to_different": float(sim[c, farthest[c]]),
}
for c in range(C)
}
return donor, picks, labels
@torch.no_grad()
def run(
a, samples, context_ids, target_ids, layers, clean, ldet, pdet, donor_map, labels, n_ctx, n_tgt
):
N = len(samples)
batches = [list(range(s, min(s + BATCH, N))) for s in range(0, N, BATCH)]
tgt_slice = list(range(n_ctx, n_ctx + n_tgt))
ti = torch.as_tensor(tgt_slice, dtype=torch.long)
rows = []
feat_map = {}
def do_condition(locus, level, layer_list):
dmap = donor_map[level]
for L in layer_list:
ca = clean["acts"][L].float()
ld = ldet[L]
feats = []
for batch in batches:
ctx = clean["context"][batch].to(DEVICE, DTYPE)
didx = dmap[batch]
donor_dev = ca[didx].to(DEVICE, DTYPE)
def hook(_m, _i, out, dd=donor_dev, loc=locus):
o = out.clone()
if loc == "target16":
o[:, ti, :] = dd[:, ti, :].to(o.dtype)
else:
o[:, :, :] = dd.to(o.dtype)
return o
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 locus == "target16":
mod[:, ti, :] = ca[didx][:, ti, :]
else:
mod = ca[didx].clone()
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": int(labels[i]),
"class_name": samples[i]["class_name"],
"human": HUMAN.get(samples[i]["class_name"], samples[i]["class_name"]),
"condition": f"{locus}_{level}",
"layer": L,
"mode": "resample",
**m,
}
)
feat_map[(f"{locus}_{level}", L, "resample")] = feats
for level in LEVELS:
do_condition("target16", level, layers)
print(f" target16 / {level}: all {len(layers)} layers done", flush=True)
for level in LEVELS:
do_condition("full", level, [6]) # layer-independent; one representative layer
print(f" full / {level}: done", flush=True)
return rows, feat_map
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)
samples = load_focus_180()
N = len(samples)
folders = sorted({s["class_name"] for s in samples}, key=lambda x: int(x))
print(f"resample-by-class: {N} samples, {len(folders)} classes", flush=True)
clean = collect_clean(a, samples, context_ids, target_ids, 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())
donor_map, picks, labels = build_donor_maps(samples, clean, folders)
print("donor class picks:", flush=True)
for f in folders:
print(
f" {f} ({HUMAN.get(f,'')}): similar->{picks[f]['similar_class']} "
f"({HUMAN.get(picks[f]['similar_class'],'')}, cos {picks[f]['sim_to_similar']:.2f}) | "
f"different->{picks[f]['different_class']} ({HUMAN.get(picks[f]['different_class'],'')})",
flush=True,
)
rows, feat_map = run(
a,
samples,
context_ids,
target_ids,
layers,
clean,
ldet,
pdet,
donor_map,
labels,
n_ctx,
n_tgt,
)
summ = aggregate(rows)
n_probes = sum(1 for s in summ if (s["condition"], s["layer"], s["mode"]) in feat_map)
trained = 0
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, str(key), use_cv=False)
trained += 1
if trained % 10 == 0 or trained == n_probes:
print(f" probe {trained}/{n_probes} trained", flush=True)
ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
rd = OUT / "ijepa_task4_deep" / "resample_by_class_180" / ts
rd.mkdir(parents=True, exist_ok=True)
(rd / "summary_metrics.json").write_text(json.dumps(summ, indent=2))
(rd / "per_sample_metrics.json").write_text(json.dumps(rows, indent=2))
(rd / "donor_class_picks.json").write_text(json.dumps(picks, indent=2))
(rd / "config.yaml").write_text(
yaml.safe_dump(
{
"run": "resample_by_class",
"n": N,
"classes": folders,
"levels": LEVELS,
"loci": ["target16", "full"],
"seed": SEED,
"maha": "ledoit-wolf",
},
sort_keys=False,
)
)
plots(rd, summ, layers)
# console summary
look = {(s["condition"], s["layer"], s["mode"]): s for s in summ}
print("\n=== target16 resample by donor class (avg over layers) ===")
for lv in LEVELS:
r = [look[(f"target16_{lv}", L, "resample")] for L in layers]
print(
f" {lv:9s}: MSE {np.mean([x['mean_prediction_mse'] for x in r]):.3f} "
f"cos {np.mean([x['mean_target_cosine'] for x in r]):.3f} "
f"probe {np.mean([x['classification_accuracy'] for x in r]):.3f}"
)
print("=== full resample by donor class (layer 6) ===")
for lv in LEVELS:
s = look[(f"full_{lv}", 6, "resample")]
print(
f" {lv:9s}: MSE {s['mean_prediction_mse']:.3f} cos {s['mean_target_cosine']:.3f} probe {s['classification_accuracy']:.3f}"
)
print(f"\nsaved -> {rd}", flush=True)
def plots(rd, summ, 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)),
("classification_accuracy", "Probe accuracy", (0, 1)),
("mean_prediction_maha", "Prediction Maha (shrinkage)", (0, None)),
]:
fig, ax = plt.subplots(figsize=(12, 5.5))
x = np.arange(len(layers))
w = 0.8 / 3
for j, lv in enumerate(LEVELS):
vals = [look.get((f"target16_{lv}", L, "resample"), {}).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=lv)
ax.set_xticks(x, [str(L) for L in layers])
ax.set_xlabel("Predictor layer")
ax.set_ylabel(ylab)
ax.set_title(f"target16 resample by donor class: {ylab}")
if ylim:
ax.set_ylim(*ylim)
ax.legend(frameon=False)
ax.grid(axis="y", alpha=0.25)
fig.tight_layout()
fig.savefig(rd / f"target16_{metric}.png", dpi=160)
plt.close(fig)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
10.6 kB
·
Xet hash:
455d9da2bbfbe97762cf30d37111de461e8b734355599256e92ec546801c48f3

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.