summaryrefslogtreecommitdiff
path: root/src/jlens_v2.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/jlens_v2.py')
-rw-r--r--src/jlens_v2.py187
1 files changed, 187 insertions, 0 deletions
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()