diff options
| -rw-r--r-- | src/jlens_v3.py | 9 | ||||
| -rw-r--r-- | src/loss_reweight.py | 197 | ||||
| -rw-r--r-- | src/synthetic_pair.py | 224 |
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) |
