summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--src/jlens_v3.py9
-rw-r--r--src/loss_reweight.py197
-rw-r--r--src/synthetic_pair.py224
3 files changed, 427 insertions, 3 deletions
diff --git a/src/jlens_v3.py b/src/jlens_v3.py
index 7316fe8..b5f59ed 100644
--- a/src/jlens_v3.py
+++ b/src/jlens_v3.py
@@ -169,6 +169,8 @@ def main():
ap.add_argument('--n_prompts', type=int, default=20)
ap.add_argument('--batch_size', type=int, default=16)
ap.add_argument('--layers', default='2,3,4', help='comma-separated layer indices')
+ ap.add_argument('--output_dir', default='outputs/jlens_v3',
+ help='directory for per-layer result artifacts')
ap.add_argument('--device', default='cuda')
ap.add_argument('--chunk', type=int, default=32)
args = ap.parse_args()
@@ -210,7 +212,7 @@ def main():
if layer_idx == n_layer - 1:
# Validation: J_{L-1} should be identity, so faithful vectors == W_U rows
sims = torch.nn.functional.cosine_similarity(
- faithful_vecs.float(), W_U.float(), dim=1)
+ faithful_vecs.float(), W_U.cpu().float(), dim=1)
print(f" [validation] last layer: mean cos-sim(faithful, W_U rows) = "
f"{sims.mean().item():.4f} (expect ~1.0 if J=identity)")
@@ -239,10 +241,11 @@ def main():
for tid, n in srt[-5:]:
print(f" '{esc(itos[tid])}' freq={freq[tid]:.3f}% norm={n:.4f}")
+ os.makedirs(args.output_dir, exist_ok=True)
torch.save({'faithful_vecs': faithful_vecs, 'faithful_norms': faithful_norms,
'proxy_norms': proxy_norms},
- f'outputs/jlens_v3_layer{layer_idx}.pt')
- print(f" Saved outputs/jlens_v3_layer{layer_idx}.pt")
+ os.path.join(args.output_dir, f'layer{layer_idx}.pt'))
+ print(f" Saved {os.path.join(args.output_dir, f'layer{layer_idx}.pt')}")
print("\nDONE.")
diff --git a/src/loss_reweight.py b/src/loss_reweight.py
new file mode 100644
index 0000000..3274807
--- /dev/null
+++ b/src/loss_reweight.py
@@ -0,0 +1,197 @@
+"""
+Loss-reweighting ablation (GPT-5.6-Terra's design, adapted to faithful J-lens).
+
+Tests whether increasing a token's EFFECTIVE frequency/importance reduces its
+faithful J-lens norm — without the confounds of the old random-insertion
+ablation (which shifted positions, destroyed n-grams, and used an unmatched
+control run).
+
+Design per seed (identical init + identical minibatch order for all three):
+ q-upweight: cross-entropy terms whose target is 'q' are weighted x2.
+ control: ordinary loss.
+ ctrl_random: same-total-loss control: weight x2 on the SAME NUMBER of
+ randomly chosen non-'q' target positions (deterministic per
+ batch index, so all models share the same control positions).
+
+If increased effective frequency causally reduces the faithful J-lens norm of
+'q', the q-upweight model must show a lower norm than BOTH controls.
+
+Steps:
+ python3 src/loss_reweight.py --step train --mode q --seed 0 [--max_iters 3000]
+ python3 src/loss_reweight.py --step jlens --mode q --seed 0 [--layers 2,3,4]
+ python3 src/loss_reweight.py --step summary
+"""
+import sys, os, argparse, subprocess, pickle
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from typing import Any
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+DATA_DIR = 'data/shakespeare_char'
+OUT_ROOT = 'out-loss-reweight'
+JLENS_OUT = 'outputs/loss_reweight'
+TARGET = 'q'
+WEIGHT = 2.0
+MODEL_ARGS: dict[str, Any] = dict(n_layer=6, n_head=6, n_embd=384, block_size=128,
+ bias=False, dropout=0.2)
+LOSS_MODES = ('q', 'control', 'ctrl_random')
+
+
+def _batch(data, blk, bs, g, device):
+ ix = torch.randint(len(data) - blk, (bs,), generator=g)
+ x = torch.stack([torch.from_numpy(data[i:i+blk].astype(np.int64)) for i in ix])
+ y = torch.stack([torch.from_numpy(data[i+1:i+1+blk].astype(np.int64)) for i in ix])
+ return x.to(device), y.to(device)
+
+
+def _weighted_loss(logits, y, mode, q_id, batch_k, V, device):
+ """Per-token weighted CE. Returns scalar loss."""
+ logp = F.log_softmax(logits.view(-1, V), dim=-1)
+ nll = -logp.gather(1, y.view(-1, 1)).squeeze(1) # (B*T,)
+ w = torch.ones_like(nll)
+ if mode == 'q':
+ w[y.view(-1) == q_id] = WEIGHT
+ elif mode == 'ctrl_random':
+ g = torch.Generator(device=device).manual_seed(1000 + batch_k)
+ n_q = int((y == q_id).sum().item())
+ flat = torch.arange(y.numel(), device=device)
+ non_q = flat[y.view(-1) != q_id]
+ if len(non_q) > 0 and n_q > 0:
+ pick = non_q[torch.randperm(len(non_q), generator=g)[:min(n_q, len(non_q))]]
+ w[pick] = WEIGHT
+ return (nll * w).mean()
+
+
+def train(mode, seed, max_iters, batch_size):
+ sys.path.insert(0, '.')
+ from model import GPT, GPTConfig
+ torch.manual_seed(seed)
+ np.random.seed(seed)
+ device = 'cuda'
+ train_data = np.memmap(f'{DATA_DIR}/train.bin', dtype=np.uint16, mode='r')
+ val_data = np.memmap(f'{DATA_DIR}/val.bin', dtype=np.uint16, mode='r')
+ with open(f'{DATA_DIR}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ q_id = meta['stoi'][TARGET]
+ args = dict(
+ n_layer=int(MODEL_ARGS['n_layer']), n_head=int(MODEL_ARGS['n_head']),
+ n_embd=int(MODEL_ARGS['n_embd']), block_size=int(MODEL_ARGS['block_size']),
+ bias=bool(MODEL_ARGS['bias']), dropout=float(MODEL_ARGS['dropout']),
+ vocab_size=int(meta['vocab_size']),
+ )
+ model = GPT(GPTConfig(**args)).to(device)
+ print(f"[{mode}] seed {seed}: params={sum(p.numel() for p in model.parameters())/1e6:.2f}M")
+
+ opt = model.configure_optimizers(weight_decay=0.1, learning_rate=1e-3,
+ betas=(0.9, 0.99), device_type='cuda')
+ bs = batch_size
+ blk = args['block_size']
+ V = args['vocab_size']
+ # identical minibatch order for every model: per-seed generator, fixed start
+ g = torch.Generator(device=device).manual_seed(20260731 + seed)
+ gval = torch.Generator(device=device).manual_seed(777 + seed)
+
+ def get_batch(split):
+ d = train_data if split == 'train' else val_data
+ gg = g if split == 'train' else gval
+ return _batch(d, blk, bs, gg, device)
+
+ best_val = 1e9
+ out_dir = f'{OUT_ROOT}/{mode}/seed{seed}'
+ os.makedirs(out_dir, exist_ok=True)
+ for it in range(max_iters):
+ if it % 500 == 0:
+ model.eval()
+ lv = []
+ for _ in range(50):
+ X, Y = get_batch('val')
+ with torch.no_grad():
+ logits = model(X)[0]
+ lv.append(F.cross_entropy(logits.view(-1, V), Y.view(-1)).item())
+ v = np.mean(lv)
+ model.train()
+ if v < best_val:
+ best_val = v
+ torch.save({'model': model.state_dict(), 'model_args': args,
+ 'best_val_loss': best_val}, f'{out_dir}/ckpt.pt')
+ if it % 1000 == 0:
+ print(f" iter {it}: val={v:.4f}")
+ X, Y = get_batch('train')
+ logits = model(X)[0]
+ loss = _weighted_loss(logits, Y, mode, q_id, it, V, device)
+ loss.backward()
+ opt.step()
+ opt.zero_grad(set_to_none=True)
+ print(f"[{mode}] seed {seed} done. best_val={best_val:.4f}")
+
+
+def jlens(mode, seed, layers, n_prompts):
+ out_dir = f'{JLENS_OUT}/{mode}/seed{seed}'
+ cmd = ["python3", "-u", "src/jlens_v3.py",
+ "--checkpoint", f'{OUT_ROOT}/{mode}/seed{seed}/ckpt.pt',
+ "--data_dir", DATA_DIR,
+ "--n_prompts", str(n_prompts),
+ "--layers", layers,
+ "--chunk", "16",
+ "--output_dir", out_dir]
+ print("running:", " ".join(cmd))
+ r = subprocess.run(cmd, cwd='/workspace/code')
+ assert r.returncode == 0, "jlens_v3 failed"
+
+
+def summary(layers):
+ with open(f'{DATA_DIR}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ q_id = meta['stoi'][TARGET]
+ seeds = sorted(set(
+ d.split('seed')[1] for mode in LOSS_MODES
+ for d in os.listdir(f'{JLENS_OUT}/{mode}')
+ if d.startswith('seed')))
+ print(f"\n{'='*78}")
+ print(f"LOSS-REWEIGHTING: faithful J-lens norm of '{TARGET}' "
+ f"({WEIGHT}x CE) vs controls")
+ print(f"{'='*78}")
+ print(f"{'seed':<5}{'layer':<6}" + "".join(f"{m:>14}" for m in LOSS_MODES))
+ for s in seeds:
+ for l in map(int, layers.split(',')):
+ row = [s, str(l)]
+ for m in LOSS_MODES:
+ d = torch.load(f'{JLENS_OUT}/{m}/seed{s}/layer{l}.pt',
+ map_location='cpu')
+ row.append(f"{d['faithful_norms'][q_id]:.4f}")
+ print(f"{row[0]:<5}{row[1]:<6}" + "".join(f"{v:>14}" for v in row[2:]))
+ # mean over middle layers per mode
+ mids = [l for l in map(int, layers.split(','))]
+ means = {}
+ for m in LOSS_MODES:
+ vals = []
+ for l in mids:
+ d = torch.load(f'{JLENS_OUT}/{m}/seed{s}/layer{l}.pt',
+ map_location='cpu')
+ vals.append(d['faithful_norms'][q_id])
+ means[m] = np.mean(vals)
+ print(f" -> mean over layers: q={means['q']:.4f} "
+ f"control={means['control']:.4f} ctrl_random={means['ctrl_random']:.4f}")
+ print(f" -> q/control = {means['q']/max(means['control'],1e-9):.3f} "
+ f"q/ctrl_random = {means['q']/max(means['ctrl_random'],1e-9):.3f}")
+
+
+if __name__ == '__main__':
+ ap = argparse.ArgumentParser()
+ ap.add_argument('--step', required=True, choices=['train', 'jlens', 'summary'])
+ ap.add_argument('--mode', choices=LOSS_MODES)
+ ap.add_argument('--seed', type=int, default=0)
+ ap.add_argument('--max_iters', type=int, default=3000)
+ ap.add_argument('--batch_size', type=int, default=16)
+ ap.add_argument('--layers', default='2,3,4')
+ ap.add_argument('--n_prompts', type=int, default=10)
+ a = ap.parse_args()
+ if a.step == 'train':
+ assert a.mode, "need --mode"
+ train(a.mode, a.seed, a.max_iters, a.batch_size)
+ elif a.step == 'jlens':
+ assert a.mode, "need --mode"
+ jlens(a.mode, a.seed, a.layers, a.n_prompts)
+ else:
+ summary(a.layers)
diff --git a/src/synthetic_pair.py b/src/synthetic_pair.py
new file mode 100644
index 0000000..7614fbe
--- /dev/null
+++ b/src/synthetic_pair.py
@@ -0,0 +1,224 @@
+"""
+Frequency-Matched Synthetic Pair Test (Gemini 3.1 Pro's design).
+
+Two synthetic tokens at IDENTICAL unigram frequency in an otherwise-normal
+Shakespeare corpus:
+ T_struct '@' : appears only after the trigger sequence "the " (high conditional
+ predictability, structured context).
+ T_noise '#' : injected at uniform random positions (zero conditional structure).
+
+Hypotheses:
+ Frequency-only: J-lens norms of '@' and '#' are identical at every layer
+ (same unigram frequency, same rarity).
+ Structure/workspace (Anthropic): '@' maintains a higher faithful J-lens norm,
+ especially in intermediate layers (model tracks the trigger
+ context in the residual stream).
+
+Uses the FAITHFUL J-lens (jlens_v3: rows of W_U * J_l) plus the old proxy.
+
+Steps:
+ python3 src/synthetic_pair.py --step prep # build data/synth_pair
+ python3 src/synthetic_pair.py --step train --seed 0 [--max_iters 3000]
+ python3 src/synthetic_pair.py --step jlens --seed 0 [--layers 0,1,2,3,4,5]
+ python3 src/synthetic_pair.py --step summary # compare @ vs #
+"""
+import sys, os, argparse, subprocess, pickle, random
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from typing import Any
+import numpy as np
+import torch
+
+DATA_SRC = 'data/shakespeare_char/input.txt'
+DATA_DIR = 'data/synth_pair'
+OUT_ROOT = 'out-synth-pair'
+JLENS_OUT = 'outputs/synth_pair'
+T_STRUCT = '@'
+T_NOISE = '#'
+TRIGGER = 'the '
+TARGET_FREQ = 0.001 # 0.1%
+MODEL_ARGS: dict[str, Any] = dict(n_layer=6, n_head=6, n_embd=384, block_size=128,
+ bias=False, dropout=0.2)
+
+
+def prep():
+ with open(DATA_SRC) as f:
+ text = f.read()
+ n_target = max(1, int(TARGET_FREQ * len(text)))
+
+ # positions for T_struct: after occurrences of TRIGGER
+ trig_positions = []
+ start = 0
+ while True:
+ i = text.find(TRIGGER, start)
+ if i < 0:
+ break
+ trig_positions.append(i + len(TRIGGER))
+ start = i + len(TRIGGER)
+ assert len(trig_positions) >= n_target, f"only {len(trig_positions)} triggers"
+ rng = np.random.RandomState(42)
+ struct_pos = sorted(rng.choice(trig_positions, size=n_target, replace=False).tolist())
+
+ # positions for T_noise: uniform random, disjoint from struct_pos
+ noise_pos = sorted(rng.choice(
+ [p for p in range(len(text)) if p not in set(struct_pos)],
+ size=n_target, replace=False).tolist())
+
+ # insert with offset (both sets sorted -> single merge pass)
+ insertions = [(p, T_STRUCT) for p in struct_pos] + [(p, T_NOISE) for p in noise_pos]
+ insertions.sort()
+ out = []
+ prev = 0
+ for pos, ch in insertions:
+ out.append(text[prev:pos])
+ out.append(ch)
+ prev = pos
+ out.append(text[prev:])
+ modified = ''.join(out)
+ assert modified.count(T_STRUCT) == modified.count(T_NOISE) == n_target
+ print(f"prep: '{T_STRUCT}' x{n_target} after '{TRIGGER.strip()}', "
+ f"'{T_NOISE}' x{n_target} random, "
+ f"freq each = {n_target/len(modified):.4%}")
+
+ # build vocab (existing chars + the two synthetic)
+ chars = sorted(set(text)) + [T_STRUCT, T_NOISE]
+ stoi = {c: i for i, c in enumerate(chars)}
+ itos = {i: c for i, c in enumerate(chars)}
+ data = np.array([stoi[c] for c in modified], dtype=np.uint16)
+ n = int(0.9 * len(data))
+ os.makedirs(DATA_DIR, exist_ok=True)
+ data[:n].tofile(os.path.join(DATA_DIR, 'train.bin'))
+ data[n:].tofile(os.path.join(DATA_DIR, 'val.bin'))
+ with open(os.path.join(DATA_DIR, 'meta.pkl'), 'wb') as f:
+ pickle.dump({'stoi': stoi, 'itos': itos, 'vocab_size': len(chars)}, f)
+ with open(os.path.join(DATA_DIR, 'input.txt'), 'w') as f:
+ f.write(modified)
+ print(f"prep: vocab={len(chars)}, train={n:,} val={len(data)-n:,} tokens")
+
+
+def train(seed, max_iters, batch_size):
+ sys.path.insert(0, '.')
+ from model import GPT, GPTConfig
+ torch.manual_seed(seed)
+ np.random.seed(seed)
+ device = 'cuda'
+ train_data = np.memmap(f'{DATA_DIR}/train.bin', dtype=np.uint16, mode='r')
+ val_data = np.memmap(f'{DATA_DIR}/val.bin', dtype=np.uint16, mode='r')
+ with open(f'{DATA_DIR}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ args = dict(
+ n_layer=int(MODEL_ARGS['n_layer']), n_head=int(MODEL_ARGS['n_head']),
+ n_embd=int(MODEL_ARGS['n_embd']), block_size=int(MODEL_ARGS['block_size']),
+ bias=bool(MODEL_ARGS['bias']), dropout=float(MODEL_ARGS['dropout']),
+ vocab_size=int(meta['vocab_size']),
+ )
+ model = GPT(GPTConfig(**args)).to(device)
+ print(f"seed {seed}: params={sum(p.numel() for p in model.parameters())/1e6:.2f}M")
+
+ opt = model.configure_optimizers(weight_decay=0.1, learning_rate=1e-3,
+ betas=(0.9, 0.99), device_type='cuda')
+ bs = batch_size
+ blk = args['block_size']
+
+ def get_batch(split):
+ d = train_data if split == 'train' else val_data
+ ix = torch.randint(len(d) - blk, (bs,))
+ x = torch.stack([torch.from_numpy(d[i:i+blk].astype(np.int64)) for i in ix])
+ y = torch.stack([torch.from_numpy(d[i+1:i+1+blk].astype(np.int64)) for i in ix])
+ return x.to(device), y.to(device)
+
+ best_val = 1e9
+ out_dir = f'{OUT_ROOT}/seed{seed}'
+ os.makedirs(out_dir, exist_ok=True)
+ for it in range(max_iters):
+ if it % 500 == 0:
+ model.eval()
+ lv = []
+ for _ in range(50):
+ X, Y = get_batch('val')
+ with torch.no_grad():
+ _, loss = model(X, Y)
+ lv.append(loss.item())
+ v = np.mean(lv)
+ model.train()
+ if v < best_val:
+ best_val = v
+ torch.save({'model': model.state_dict(), 'model_args': args,
+ 'best_val_loss': best_val}, f'{out_dir}/ckpt.pt')
+ if it % 1000 == 0:
+ print(f" iter {it}: val={v:.4f}")
+ X, Y = get_batch('train')
+ _, loss = model(X, Y)
+ loss.backward()
+ opt.step()
+ opt.zero_grad(set_to_none=True)
+ print(f"seed {seed} done. best_val={best_val:.4f}")
+
+
+def jlens(seed, layers, n_prompts):
+ out_dir = f'{JLENS_OUT}/seed{seed}'
+ cmd = ["python3", "-u", "src/jlens_v3.py",
+ "--checkpoint", f'{OUT_ROOT}/seed{seed}/ckpt.pt',
+ "--data_dir", DATA_DIR,
+ "--n_prompts", str(n_prompts),
+ "--layers", layers,
+ "--chunk", "16",
+ "--output_dir", out_dir]
+ print("running:", " ".join(cmd))
+ r = subprocess.run(cmd, cwd='/workspace/code')
+ assert r.returncode == 0, "jlens_v3 failed"
+
+
+def summary(layers):
+ with open(f'{DATA_DIR}/meta.pkl', 'rb') as f:
+ meta = pickle.load(f)
+ sid = meta['stoi']
+ iid_s = sid[T_STRUCT]
+ iid_n = sid[T_NOISE]
+ train_data = np.memmap(f'{DATA_DIR}/train.bin', dtype=np.uint16, mode='r')
+ V = meta['vocab_size']
+ counts = np.bincount(train_data, minlength=V).astype(float)
+ freq = counts / counts.sum() * 100
+
+ seeds = sorted([d for d in os.listdir(JLENS_OUT) if d.startswith('seed')])
+ print(f"\n{'='*72}")
+ print("SYNTHETIC PAIR: faithful J-lens norms of T_struct '@' vs T_noise '#'")
+ print(f"('@' freq={freq[iid_s]:.3f}%, '#' freq={freq[iid_n]:.3f}%)")
+ print(f"{'='*72}")
+ print(f"{'seed':<5}{'layer':<6}{'@ norm':>9}{'# norm':>9}{'ratio':>8} freq-corr r")
+ for s in seeds:
+ for l in map(int, layers.split(',')):
+ d = torch.load(f'{JLENS_OUT}/{s}/layer{l}.pt', map_location='cpu')
+ fn = d['faithful_norms']
+ f_arr = np.array([freq[k] for k in range(V)])
+ rf = np.corrcoef(np.array([fn[k] for k in range(V)]), f_arr)[0, 1]
+ print(f"{s:<5}{l:<6}{fn[iid_s]:>9.4f}{fn[iid_n]:>9.4f}"
+ f"{fn[iid_s]/max(fn[iid_n],1e-9):>8.2f} {rf:+.3f}")
+ # per-seed ratio across middle layers
+ mids = [l for l in map(int, layers.split(',')) if l in (2, 3, 4)]
+ rs = []
+ rn = []
+ for l in mids:
+ d = torch.load(f'{JLENS_OUT}/{s}/layer{l}.pt', map_location='cpu')
+ fn = d['faithful_norms']
+ rs.append(fn[iid_s])
+ rn.append(fn[iid_n])
+ print(f" -> mean middle-layer ratio @/# = {np.mean(rs)/max(np.mean(rn),1e-9):.3f}")
+
+
+if __name__ == '__main__':
+ ap = argparse.ArgumentParser()
+ ap.add_argument('--step', required=True, choices=['prep', 'train', 'jlens', 'summary'])
+ ap.add_argument('--seed', type=int, default=0)
+ ap.add_argument('--max_iters', type=int, default=3000)
+ ap.add_argument('--batch_size', type=int, default=16)
+ ap.add_argument('--layers', default='0,1,2,3,4,5')
+ ap.add_argument('--n_prompts', type=int, default=10)
+ a = ap.parse_args()
+ if a.step == 'prep':
+ prep()
+ elif a.step == 'train':
+ train(a.seed, a.max_iters, a.batch_size)
+ elif a.step == 'jlens':
+ jlens(a.seed, a.layers, a.n_prompts)
+ else:
+ summary(a.layers)