Buckets:
| #!/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 | |
| 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.