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