summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--src/jlens.py19
1 files 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")