File size: 5,054 Bytes
f6c2fe1 | 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 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 | import os, json, math, sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from configs.paths import dim_paths
from src.utils import read_json, write_json
DIM = "monitoring"
p = dim_paths(DIM)
ACT_PATH = p.ACTIVATIONS
FULL_PATH = p.DIRECTIONS
NO_ORTHO_PATH = os.path.join(p.CHECKPOINT_DIR, "directions_monitoring_noOrtho.pt")
NO_PCA_PATH = os.path.join(p.CHECKPOINT_DIR, "directions_monitoring_noPCA.pt")
SEL_PATH = os.path.join(p.CHECKPOINT_DIR, "selected_layers_monitoring_allmonoV2.json")
def flat(v):
if isinstance(v, dict):
v = v.get("direction", v.get("vec", v.get("vector", v)))
if isinstance(v, torch.Tensor):
v = v.detach().float()
else:
v = torch.tensor(v).float()
return v.view(-1)
def cos(a, b):
a = flat(a)
b = flat(b)
if a.numel() != b.numel():
return None
an = a.norm()
bn = b.norm()
if an < 1e-8 or bn < 1e-8:
return None
return float(torch.dot(a, b) / (an * bn))
def get_dir(blob, L):
d = blob["directions"]
if L in d:
return flat(d[L])
if str(L) in d:
return flat(d[str(L)])
return None
acts_blob = torch.load(ACT_PATH, map_location="cpu", weights_only=False)
full_blob = torch.load(FULL_PATH, map_location="cpu", weights_only=False)
no_ortho_blob = torch.load(NO_ORTHO_PATH, map_location="cpu", weights_only=False)
no_pca_blob = torch.load(NO_PCA_PATH, map_location="cpu", weights_only=False)
sel = read_json(SEL_PATH)
selected = set(int(x) for x in sel["selected_layers"])
rows = []
for L_raw, data in acts_blob["per_layer"].items():
L = int(L_raw)
acts = data["acts"].float()
labels = data["labels"]
pos = acts[labels == 1]
neg = acts[labels == 0]
if pos.shape[0] < 5 or neg.shape[0] < 5:
continue
raw_md = pos.mean(0) - neg.mean(0)
mu = acts.mean(0)
d_full = get_dir(full_blob, L)
d_no_ortho = get_dir(no_ortho_blob, L)
d_no_pca = get_dir(no_pca_blob, L)
row = {
"layer": L,
"selected": L in selected,
"n_pos": int(pos.shape[0]),
"n_neg": int(neg.shape[0]),
"raw_norm": float(raw_md.norm()),
"mu_norm": float(mu.norm()),
"cos_raw_mu": cos(raw_md, mu),
"cos_noOrtho_mu": cos(d_no_ortho, mu) if d_no_ortho is not None else None,
"cos_noPCA_mu": cos(d_no_pca, mu) if d_no_pca is not None else None,
"cos_full_mu": cos(d_full, mu) if d_full is not None else None,
"abs_cos_raw_mu": abs(cos(raw_md, mu)) if cos(raw_md, mu) is not None else None,
"abs_cos_noOrtho_mu": abs(cos(d_no_ortho, mu)) if d_no_ortho is not None and cos(d_no_ortho, mu) is not None else None,
"abs_cos_noPCA_mu": abs(cos(d_no_pca, mu)) if d_no_pca is not None and cos(d_no_pca, mu) is not None else None,
"abs_cos_full_mu": abs(cos(d_full, mu)) if d_full is not None and cos(d_full, mu) is not None else None,
}
if row["abs_cos_noOrtho_mu"] is not None and row["abs_cos_full_mu"] is not None:
row["ortho_overlap_reduction_mu"] = row["abs_cos_noOrtho_mu"] - row["abs_cos_full_mu"]
else:
row["ortho_overlap_reduction_mu"] = None
rows.append(row)
def avg(xs):
xs = [x for x in xs if x is not None]
return sum(xs) / len(xs) if xs else None
def group_summary(name, rs):
return {
"group": name,
"n_layers": len(rs),
"mean_abs_cos_raw_mu": avg([r["abs_cos_raw_mu"] for r in rs]),
"mean_abs_cos_noOrtho_mu": avg([r["abs_cos_noOrtho_mu"] for r in rs]),
"mean_abs_cos_noPCA_mu": avg([r["abs_cos_noPCA_mu"] for r in rs]),
"mean_abs_cos_full_mu": avg([r["abs_cos_full_mu"] for r in rs]),
"mean_ortho_overlap_reduction_mu": avg([r["ortho_overlap_reduction_mu"] for r in rs]),
}
selected_rows = [r for r in rows if r["selected"]]
rejected_rows = [r for r in rows if not r["selected"]]
summary = {
"note": "Geometry diagnostic from cached contrastive activations. cos_mu tests overlap with general reasoning mean.",
"activation_path": ACT_PATH,
"full_direction_path": FULL_PATH,
"no_ortho_path": NO_ORTHO_PATH,
"no_pca_path": NO_PCA_PATH,
"selected_layer_file": SEL_PATH,
"groups": [
group_summary("selected_layers", selected_rows),
group_summary("rejected_or_unselected_layers", rejected_rows),
group_summary("all_layers", rows),
],
"rows": rows,
}
out = os.path.join(p.RESULTS_DIR, "orthogonalization_geometry_mu_summary.json")
write_json(summary, out)
print("Saved:", out)
print()
print("| group | n | raw cos μ | noOrtho cos μ | noPCA cos μ | full cos μ | ortho reduction |")
print("|---|---:|---:|---:|---:|---:|---:|")
for g in summary["groups"]:
print(
f"| {g['group']} | {g['n_layers']} "
f"| {g['mean_abs_cos_raw_mu']:.4f} "
f"| {g['mean_abs_cos_noOrtho_mu']:.4f} "
f"| {g['mean_abs_cos_noPCA_mu']:.4f} "
f"| {g['mean_abs_cos_full_mu']:.4f} "
f"| {g['mean_ortho_overlap_reduction_mu']:.4f} |"
)
|