summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-07-30 15:43:40 +0100
committerVoid Agent <void@jayrup.hermes>2026-07-30 15:43:40 +0100
commit763563c775660bc77c8866abe9ecf7c1764c3a71 (patch)
tree2e5c9f8d2ef1de29d3f3a3295371e2f8f376d40c
parent86b6b16c4473e11681ceeeea03a57cd98eaf24e1 (diff)
Add Pythia-70m test script
-rw-r--r--src/test_pythia.py51
1 files changed, 51 insertions, 0 deletions
diff --git a/src/test_pythia.py b/src/test_pythia.py
new file mode 100644
index 0000000..4a05192
--- /dev/null
+++ b/src/test_pythia.py
@@ -0,0 +1,51 @@
+"""
+Test loading Pythia-70m on K2200.
+Usage: python3 src/test_pythia.py
+"""
+import sys, os
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+import torch, time, subprocess
+
+# Install if needed
+try:
+ import transformers
+except ImportError:
+ print("Installing transformers...")
+ subprocess.check_call([sys.executable, '-m', 'pip', 'install', '-q', 'transformers', 'huggingface_hub'])
+ import transformers
+
+from transformers import GPTNeoXForCausalLM
+
+model_name = 'EleutherAI/pythia-70m'
+print(f"Loading {model_name} (step 1000 checkpoint)...")
+t0 = time.time()
+
+model = GPTNeoXForCausalLM.from_pretrained(
+ model_name,
+ revision='step1000',
+ torch_dtype=torch.float32,
+ device_map='cuda'
+)
+model.eval()
+
+mem = torch.cuda.max_memory_allocated() / 1e9
+total = torch.cuda.get_device_properties(0).total_memory / 1e9
+print(f"Loaded in {time.time()-t0:.1f}s")
+print(f"VRAM: {mem:.1f}GB / {total:.1f}GB")
+print(f"Params: {sum(p.numel() for p in model.parameters())/1e6:.1f}M")
+print(f"Layers: {len(model.gpt_neox.layers)}")
+print(f"Hidden: {model.config.hidden_size}")
+print(f"Vocab: {model.config.vocab_size}")
+
+# Quick forward pass test
+tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
+inputs = tokenizer("Hello world", return_tensors="pt").to('cuda')
+with torch.no_grad():
+ outputs = model(**inputs)
+print(f"Forward pass OK, logits shape: {outputs.logits.shape}")
+
+# Check available checkpoints
+print("\nPythia-70m checkpoint revisions available:")
+print(" step1 through step143000 (154 total)")
+print(" Key steps: 1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1K, 2K, ..., 143K")
+print("\nSUCCESS — Pythia-70m fits comfortably on K2200!")