#!/usr/bin/env python3
"""Inspect what's actually in the checkpoint.pt file"""
import torch
from pathlib import Path

ckpt_path = Path("storage/proto_lm/checkpoint.pt")

if not ckpt_path.exists():
    print(f"ERROR: {ckpt_path} does not exist!")
    exit(1)

print(f"Loading checkpoint from: {ckpt_path}")
print(f"File size: {ckpt_path.stat().st_size / 1024:.1f} KB")
print()

try:
    ckpt = torch.load(ckpt_path, map_location='cpu')
    
    print("Checkpoint keys:", list(ckpt.keys()))
    print()
    
    if 'step' in ckpt:
        print(f"Step: {ckpt['step']}")
    
    if 'vocab' in ckpt:
        vocab_data = ckpt['vocab']
        if isinstance(vocab_data, dict):
            tokens = vocab_data.get('tokens', [])
            merges = vocab_data.get('merges', [])
            print(f"Vocab size: {len(tokens)}")
            print(f"Merges: {len(merges)}")
            print(f"Tokens [110:120]: {tokens[110:120]}")
        else:
            print(f"Vocab type: {type(vocab_data)}")
    
    if '_metadata' in ckpt:
        print(f"Metadata: {ckpt['_metadata']}")
    
    if 'model' in ckpt:
        model_state = ckpt['model']
        embed_weight = model_state.get('embed.weight')
        if embed_weight is not None:
            print(f"Embedding shape: {embed_weight.shape}")
            print(f"Embedding dtype: {embed_weight.dtype}")
    
except Exception as e:
    print(f"ERROR loading checkpoint: {e}")
    import traceback
    traceback.print_exc()
