From f78b0b53b821b2ecc140efa980924255b359460a Mon Sep 17 00:00:00 2001 From: Void Agent Date: Wed, 29 Jul 2026 19:44:47 +0100 Subject: jlens.py: fix load_model for dict-format nanoGPT checkpoints --- src/jlens.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/src/jlens.py b/src/jlens.py index f332d44..0cff809 100644 --- a/src/jlens.py +++ b/src/jlens.py @@ -26,14 +26,21 @@ from pathlib import Path from collections import defaultdict -def load_model(checkpoint_path, model_class, device='cuda'): +def load_model(checkpoint_path, device='cuda'): """Load a trained nanoGPT model from checkpoint.""" checkpoint = torch.load(checkpoint_path, map_location=device) - # nanoGPT stores model args, state_dict + optimizer in checkpoint + # nanoGPT stores model args as a dict in checkpoint model_args = checkpoint['model_args'] - # Create model with saved config - model = model_class(model_args) + # Convert dict to GPTConfig if needed + if isinstance(model_args, dict): + from model import GPTConfig + config = GPTConfig(**model_args) + else: + config = model_args + + # Create model with config + model = GPT(config) # Fix state dict keys (nanoGPT wraps in DataParallel) state_dict = checkpoint['model'] @@ -45,7 +52,7 @@ def load_model(checkpoint_path, model_class, device='cuda'): model.load_state_dict(state_dict) model.to(device) model.eval() - return model, model_args + return model, config def compute_jlens_single_token(model, token_id, dataloader, layer_idx, device='cuda'): @@ -355,7 +362,7 @@ if __name__ == '__main__': # Load model print(f"Loading model from {args.checkpoint}") - model, model_args = load_model(args.checkpoint, GPT, args.device) + model, model_args = load_model(args.checkpoint, args.device) print(f"Model: {model_args.n_layer} layers, {model_args.n_embd} dim, " f"{model_args.n_head} heads, {model_args.vocab_size} vocab") -- cgit v1.2.3