From e4ce11b38f58a562c09cb9d28c37358058ff0225 Mon Sep 17 00:00:00 2001 From: Void Agent Date: Wed, 29 Jul 2026 19:47:47 +0100 Subject: Add jlens_v2.py: simplified gradient approach for J-lens --- src/jlens_v2.py | 187 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 187 insertions(+) create mode 100644 src/jlens_v2.py diff --git a/src/jlens_v2.py b/src/jlens_v2.py new file mode 100644 index 0000000..1bacbc6 --- /dev/null +++ b/src/jlens_v2.py @@ -0,0 +1,187 @@ +""" +J-lens v2: simpler gradient approach that works reliably. +Uses register_full_backward_hook to capture gradients from a dummy loss. +""" +import torch +import torch.nn as nn +import numpy as np +import pickle +import os +from collections import defaultdict + +def load_model(checkpoint_path, device='cuda'): + checkpoint = torch.load(checkpoint_path, map_location=device) + model_args = checkpoint['model_args'] + from model import GPT, GPTConfig + config = GPTConfig(**model_args) + model = GPT(config) + state_dict = checkpoint['model'] + unwanted = '_orig_mod.' + for k in list(state_dict.keys()): + if k.startswith(unwanted): + state_dict[k[len(unwanted):]] = state_dict.pop(k) + model.load_state_dict(state_dict) + model.to(device) + model.eval() + return model, config + + +def compute_jlens_layer(model, layer_idx, data, device='cuda'): + """ + Compute J-lens for all tokens at one layer. + + For each token k in vocab, J-lens = E_x[ grad of log p(k|x) w.r.t. residual at layer ] + + Efficient approach: use torch.autograd.grad on log_softmax(logits)[:,:,k].sum() + """ + n_embd = model.config.n_embd + vocab_size = model.config.vocab_size + + # Accumulators + accum = torch.zeros(vocab_size, n_embd, device='cpu') + counts = torch.zeros(vocab_size, device='cpu') + + print(f" Layer {layer_idx}: processing {len(data)} batches...") + + for batch_idx, (x, y) in enumerate(data): + x, y = x.to(device), y.to(device) + B, T = x.shape + + # Store intermediate activations by splitting the forward pass + # Run model up to the target layer, capture output, + # then run the rest + + # Approach: forward hook captures block output + resid_captured = {} + def hook(module, input, output): + resid_captured['val'] = output + + target_block = model.transformer.h[layer_idx] + handle = target_block.register_forward_hook(hook) + + # Full forward pass + logits, loss = model(x, y) + handle.remove() + + # log_softmax for per-token log-probs + log_probs = torch.nn.functional.log_softmax(logits, dim=-1) # (B, T, V) + + # For each token, compute gradient w.r.t. captured residual + for token_id in range(vocab_size): + token_log_prob = log_probs[:, :, token_id].sum() + + grad = torch.autograd.grad( + token_log_prob, + resid_captured['val'], + retain_graph=(token_id < vocab_size - 1) + )[0] # (B, T, n_embd) + + # Sum over batch and sequence positions, accumulate on CPU + accum[token_id] += grad.detach().cpu().reshape(-1, n_embd).sum(dim=0) + counts[token_id] += B * T + + # Cleanup + del logits, loss, log_probs + del grad # noqa: F821 + resid_captured.clear() + torch.cuda.empty_cache() + + if (batch_idx + 1) % 5 == 0: + print(f" Batch {batch_idx + 1}/{len(data)} done") + + # Average + jlens = {} + for token_id in range(vocab_size): + if counts[token_id] > 0: + jlens[token_id] = accum[token_id] / counts[token_id] + else: + jlens[token_id] = torch.zeros(n_embd) + + return jlens + + +def main(): + import argparse + parser = argparse.ArgumentParser() + parser.add_argument('--checkpoint', type=str, required=True) + parser.add_argument('--data_dir', type=str, default='data/shakespeare_char') + parser.add_argument('--output_dir', type=str, default='outputs/jlens') + parser.add_argument('--max_batches', type=int, default=20) + parser.add_argument('--batch_size', type=int, default=16) + parser.add_argument('--device', type=str, default='cuda') + args = parser.parse_args() + + print(f"Loading model from {args.checkpoint}") + model, config = load_model(args.checkpoint, args.device) + print(f"Model: {config.n_layer} layers, {config.n_embd} dim, " + f"{config.n_head} heads, {config.vocab_size} vocab") + + # Load data + from pathlib import Path + data_dir = Path(args.data_dir) + train_data = np.memmap(data_dir / 'train.bin', dtype=np.uint16, mode='r') + meta_path = data_dir / 'meta.pkl' + with open(meta_path, 'rb') as f: + meta = pickle.load(f) + itos = meta['itos'] + print(f"Vocabulary size: {len(itos)}") + + # Build batched data + block_size = config.block_size + batches = [] + for _ in range(args.max_batches): + ix = torch.randint(len(train_data) - block_size, (args.batch_size,)) + x = torch.stack([torch.from_numpy( + train_data[i:i+block_size].astype(np.int64)) for i in ix]) + y = torch.stack([torch.from_numpy( + train_data[i+1:i+1+block_size].astype(np.int64)) for i in ix]) + batches.append((x, y)) + + print(f"Processing {len(batches)} batches of size {args.batch_size}") + + # Compute J-lens for all layers + all_jlens = {} + for layer_idx in range(config.n_layer): + print(f"\n=== Layer {layer_idx}/{config.n_layer} ===") + all_jlens[layer_idx] = compute_jlens_layer(model, layer_idx, batches, args.device) + + # Save + os.makedirs(args.output_dir, exist_ok=True) + save_path = os.path.join(args.output_dir, 'jlens_vectors.pkl') + output = { + 'metadata': { + 'n_layers': config.n_layer, + 'n_embd': config.n_embd, + 'vocab_size': config.vocab_size, + 'num_batches': len(batches), + 'batch_size': args.batch_size, + }, + 'vectors': { + str(layer): {str(tid): v.numpy() for tid, v in tokens.items()} + for layer, tokens in all_jlens.items() + } + } + with open(save_path, 'wb') as f: + pickle.dump(output, f) + + # Quick analysis + print(f"\n{'='*60}") + print("J-LENS ANALYSIS") + print(f"{'='*60}") + for layer_idx in range(config.n_layer): + norms = {tid: v.norm().item() for tid, v in all_jlens[layer_idx].items()} + sorted_tokens = sorted(norms.items(), key=lambda x: x[1], reverse=True) + print(f"\nLayer {layer_idx} — Top 10 tokens:") + for tid, norm in sorted_tokens[:10]: + tok = itos[tid].replace('\n', '\\n') + print(f" '{tok}': norm={norm:.4f}") + + threshold = np.median(list(norms.values())) * 2 + active = sum(1 for n in norms.values() if n > threshold) + print(f" Active tokens (norm > {threshold:.2f}): {active}/{len(norms)}") + + print(f"\nSaved to {save_path}") + + +if __name__ == '__main__': + main() -- cgit v1.2.3