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()