File size: 7,127 Bytes
0cd77fb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 | """Does the model use SPATIAL CONTENT, or just a constant bias on the bank pathway?
Zero-ablation cannot tell these apart: if the banks carried a learned constant, zeroing it would be
equally catastrophic. This isolates the question. CPU-only -- never touches the training GPUs.
T1 BANK DIVERSITY Are the banks actually different per sample? If banks are ~identical across
different observations, they ARE a constant and nothing else matters.
T2 BUILDER SENSITIVITY Feed sample A's DA3 features vs sample B's -- do the banks change?
T3 INTERVENTION LOSS Same batch, same rng, four conditions:
correct / shuffled (each sample gets ANOTHER's geometry) / mean-bank
(per-sample info removed, magnitude kept) / zeroed.
shuffled ~ correct => model ignores WHICH geometry (generic signal)
shuffled ~ zeroed => model reads per-sample geometry (what we want)
mean ~ correct => a constant component carries everything
"""
import os, sys, dataclasses, json
import numpy as np, jax, jax.numpy as jnp, flax.nnx as nnx
sys.path.insert(0, "scripts")
CKPT = os.environ["CONTRIB_CKPT"]
BS = int(os.environ.get("SPEC_BS", "4"))
NFS = int(os.environ.get("SPEC_NFS", "4"))
def main():
os.environ.update(USE_DA3_FULL="1", DA3_INIT_STD="0.01", DA3_LOGIT_GAIN="1",
B1K_ACTIVITIES="clean_up_your_desk", B1K_EXTRACT_DEVICES=os.environ.get("B1K_EXTRACT_DEVICES","cpu"),
B1K_2026_ROOT="/work/jack/behavior1k/data/behavior_2026_task29_42",
B1K_INIT_PARAMS=CKPT)
import train_2026 as t26
from b1k.training.b1k_da3 import create_v3_behavior_da3_loader
config = dataclasses.replace(t26.build_config(), batch_size=BS, num_workers=0)
sharding = jax.sharding.SingleDeviceSharding(jax.devices("cpu")[0])
print(f"[1] building CPU loader (BS={BS})...", flush=True)
loader = create_v3_behavior_da3_loader(
config, os.environ["B1K_2026_ROOT"], ["clean_up_your_desk"],
"/work/jack/behavior1k/task_data.json",
lang_cache="/work/jack/behavior1k/modernbert_b1k_tasks.pkl",
sharding=sharding, shuffle=True, num_workers=0)
obs, actions = next(iter(loader))
model = config.model.create(jax.random.key(0))
loaded = config.weight_loader.load(jax.tree.map(np.asarray, nnx.state(model).to_pure_dict()))
gd, st = nnx.split(model); st.replace_by_pure_dict(loaded); model = nnx.merge(gd, st)
raw = json.load(open("/work/jack/behavior1k/checkpoints/behavior_50t_checkpoint/assets/"
"IliaLarchenko/behavior_224_rgb/norm_stats.json"))["norm_stats"]
model.load_correlation_matrix({"actions": {k: (np.asarray(v, np.float32) if isinstance(v, list) else v)
for k, v in raw["actions"].items()}})
print("[2] checkpoint loaded", flush=True)
banks = model._compute_banks(obs)
keys = sorted(banks.keys()) if isinstance(banks, dict) else None
print(f"[3] banks type={type(banks).__name__} keys={keys}", flush=True)
# ---------- T1: are banks different per sample? ----------
print("\n=== T1 BANK DIVERSITY (constant check) ===")
print(" cos~1.000 between different samples => the bank IS effectively a constant")
for k in keys:
b = np.asarray(banks[k], np.float32) # [B, tokens, dim]
flat = b.reshape(b.shape[0], -1)
n = flat / (np.linalg.norm(flat, axis=1, keepdims=True) + 1e-9)
C = n @ n.T
off = C[~np.eye(C.shape[0], dtype=bool)]
# variance across samples vs across tokens: if per-sample var << per-token var, it's constant
var_across_samples = float(b.var(axis=0).mean())
var_across_tokens = float(b.var(axis=1).mean())
mean_b = b.mean(axis=0, keepdims=True)
resid = float(np.linalg.norm(b - mean_b) / (np.linalg.norm(b) + 1e-9))
print(f" {k:5s} shape={b.shape} pairwise cos across samples: mean={off.mean():+.4f} "
f"max={off.max():+.4f}")
print(f" var(across samples)={var_across_samples:.5f} var(across tokens)={var_across_tokens:.5f}"
f" ||b-mean||/||b||={resid:.4f} (near 0 => constant)")
# ---------- T2: is the bank builder sensitive to its DA3 input? ----------
print("\n=== T2 BUILDER INPUT SENSITIVITY ===")
print(" swap sample0's DA3 features for sample1's; banks must change if geometry is read")
feats = obs.da3_features
swapped = feats.at[0].set(feats[1])
obs_sw = dataclasses.replace(obs, da3_features=swapped)
banks_sw = model._compute_banks(obs_sw)
for k in keys:
a = np.asarray(banks[k], np.float32)[0]
c = np.asarray(banks_sw[k], np.float32)[0]
rel = float(np.linalg.norm(c - a) / (np.linalg.norm(a) + 1e-9))
print(f" {k:5s} rel change of sample0's bank after feature swap = {rel:.4f} "
f"{'<-- INSENSITIVE (bank ignores DA3 features!)' if rel < 1e-3 else ''}")
# ---------- T3: loss under interventions ----------
print("\n=== T3 INTERVENTION LOSS (same batch, same rng) ===")
rng = jax.random.key(7)
orig_fn = model._compute_banks
def run(tag, make):
# new _compute_banks signature: (obs, return_aux=False, depth_drop_rng=None) -> banks
# or (banks, aux) when return_aux. compute_detailed_loss calls it with return_aux=True.
model._compute_banks = (lambda o, return_aux=False, depth_drop_rng=None:
((make(), None) if return_aux else make()))
ld = model.compute_detailed_loss(rng, obs, actions, train=False, num_flow_samples=NFS)
model._compute_banks = orig_fn
a = float(jnp.mean(ld["action_loss"]))
print(f" {tag:22s} action_loss={a:.5f}", flush=True)
return a
base = run("correct", lambda: banks)
shuf = run("shuffled (roll +1)", lambda: {k: jnp.roll(banks[k], 1, axis=0) for k in keys})
meanb = run("mean-bank", lambda: {k: jnp.broadcast_to(banks[k].mean(axis=0, keepdims=True),
banks[k].shape) for k in keys})
zero = run("zeroed", lambda: None)
print("\n=== VERDICT ===")
d_sh, d_mn, d_z = shuf - base, meanb - base, zero - base
print(f" shuffled delta = {d_sh:+.5f} ({100*d_sh/max(base,1e-6):+.1f}%)")
print(f" mean-bank delta = {d_mn:+.5f} ({100*d_mn/max(base,1e-6):+.1f}%)")
print(f" zeroed delta = {d_z:+.5f} ({100*d_z/max(base,1e-6):+.1f}%)")
frac = d_sh / max(d_z, 1e-9)
print(f" shuffled/zeroed damage ratio = {frac:.3f}")
if d_sh < 0.05 * max(d_z, 1e-9):
print(" => banks act as a GENERIC/CONSTANT signal: wrong geometry costs almost nothing.")
elif frac > 0.3:
print(" => model reads PER-SAMPLE SPATIAL CONTENT: wrong geometry is nearly as bad as none.")
else:
print(" => partial: some per-sample use, but a large generic component.")
print("SPECIFICITY CHECK DONE", flush=True)
if __name__ == "__main__":
main()
|