"""Experiment 3: J-space occupancy. Decompose residual activations as sparse non-negative combinations of k<=25 J-lens vectors (matching pursuit + NNLS refit) and measure the fraction of activation variance the J-space carries per layer (paper: ~10%, intermediate layers only). """ import os, sys from pathlib import Path import numpy as np import torch from scipy.optimize import nnls sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from jlens.core import JLens, chat_ids, collect_residuals, load_model RES = Path(os.environ.get("JLENS_RESULTS", "results")) K = 25 PROMPTS = [ "The animal that spins webs has how many legs? Answer with just a number.", "Write one sentence about the ocean.", "What is the capital of France? Answer with just the city name.", "Explain photosynthesis in one sentence.", ] def pursuit_r2(D, h, k=K): """Non-negative matching pursuit of h onto rows of D. Returns R^2, support.""" r = h.clone() sel = [] for _ in range(k): scores = D @ r if sel: scores[torch.tensor(sel, device=D.device)] = -1e30 j = int(scores.argmax()) if scores[j] <= 0: break sel.append(j) A = D[sel].T # (d, |sel|) x, _ = nnls(A.detach().cpu().numpy().astype(np.float64), h.detach().cpu().numpy().astype(np.float64)) approx = A @ torch.tensor(x, device=D.device, dtype=D.dtype) r = h - approx r2 = 1 - (r.norm() / h.norm()) ** 2 return float(r2), sel def main(): model, tok = load_model(dtype=torch.bfloat16) ck = torch.load(sys.argv[1] if len(sys.argv) > 1 else "results/jbar.pt", map_location="cuda") jl = JLens(model, tok, ck["Jbar"].cuda()) WU = model.lm_head.weight.detach().float() # (V, d) L = len(model.model.language_model.layers) layers = list(range(1, L, 3)) r2_by_layer = {l: [] for l in layers} for p in PROMPTS: ids = chat_ids(tok, p) hs = collect_residuals(model, ids) T = ids.shape[1] positions = list(range(max(1, T - 12), T)) # skip BOS region for l in layers: D = WU @ jl.Jbar[l] # (V, d) J-lens dictionary at layer l D = D / D.norm(dim=1, keepdim=True).clamp_min(1e-8) for t in positions: r2, _ = pursuit_r2(D, hs[l, t].float()) r2_by_layer[l].append(r2) del D torch.cuda.empty_cache() print(f"done prompt: {p[:40]}...", flush=True) print("\nlayer | mean R^2 of k<=25 non-negative J-lens pursuit") means = {} for l in layers: means[l] = float(np.mean(r2_by_layer[l])) print(f"L{l:>2} | {means[l]:.3f}") torch.save(means, RES / "jspace_r2.pt") if __name__ == "__main__": main()