| |
| """Fresh eval for CompactLM-5M: val perplexity + multi-prompt generation + degeneracy check. |
| Imports the model class from train_compactlm5m.py so we load the EXACT architecture. |
| """ |
| import os, sys, math, json, re |
| import numpy as np |
| import torch |
|
|
| sys.path.insert(0, "/work") |
| from train_compactlm5m import CompactLM, load_tok, CTX |
|
|
| OUT = "/work/models/compactlm-5m" |
| TOK = load_tok() |
| VOCAB = 12288 |
|
|
| |
| ck = torch.load(os.path.join(OUT, "best.pt"), map_location="cpu", weights_only=False) |
| model = CompactLM(vocab=VOCAB, d=256, n_layers=4, n_heads=4, ff=640, ctx=CTX) |
| sd = ck["model"] |
| if hasattr(sd, "state_dict"): |
| sd = sd.state_dict() |
| missing, unexpected = model.load_state_dict(sd, strict=False) |
| print("missing:", missing) |
| print("unexpected:", unexpected) |
| model.eval() |
|
|
| print("n_params:", sum(p.numel() for p in model.parameters())) |
|
|
| |
| val_npy = os.path.join(OUT, "data", "val.npy") |
| if os.path.exists(val_npy): |
| val = np.load(val_npy) |
| |
| n_win = min(256, len(val)) |
| idx = torch.from_numpy(val[:n_win]).long() |
| with torch.no_grad(): |
| logits = model(idx) |
| loss = torch.nn.functional.cross_entropy( |
| logits[:, :-1].reshape(-1, VOCAB).float(), |
| idx[:, 1:].reshape(-1), ignore_index=-1) |
| ppl = math.exp(loss.item()) |
| print(f"VAL: loss={loss.item():.4f} ppl={ppl:.2f} over {n_win*CTX:,} tok") |
| else: |
| print("no val.npy") |
|
|
| |
| prompts = [ |
| "The cat sat on the", |
| "Once upon a time", |
| "The sun rises in the", |
| "I like to eat", |
| "Water boils at", |
| ] |
| results = [] |
| for p in prompts: |
| ids = torch.tensor([TOK.encode(p, add_special_tokens=False).ids]) |
| for seed in [0, 1, 2]: |
| out = model.generate(ids, max_new_tokens=64, temperature=0.8, top_k=40, seed=seed) |
| text = TOK.decode(out[0].tolist(), skip_special_tokens=True) |
| results.append({"prompt": p, "seed": seed, "text": text}) |
| print(f"\n=== {p!r} seed={seed} ===\n{text}") |
|
|
| |
| def degenerate(text): |
| |
| words = text.split() |
| if len(words) < 6: |
| return False, "short" |
| |
| tail = words[-40:] |
| seen = {} |
| rep = 0 |
| for i in range(len(tail) - 2): |
| g = tuple(tail[i:i+3]) |
| seen[g] = seen.get(g, 0) + 1 |
| maxrep = max(seen.values()) |
| frac = maxrep * 3 / len(tail) |
| return frac > 0.6, f"max3gram_frac={frac:.2f}" |
|
|
| degen_count = 0 |
| for r in results: |
| d, why = degenerate(r["text"]) |
| r["degenerate"] = d |
| r["why"] = why |
| if d: |
| degen_count += 1 |
|
|
| print(f"\nDEGENERACY: {degen_count}/{len(results)} degenerate") |
| with open(os.path.join(OUT, "eval_fresh.json"), "w") as f: |
| json.dump({"val_ppl": ppl if os.path.exists(val_npy) else None, |
| "val_loss": loss.item() if os.path.exists(val_npy) else None, |
| "samples": results, "degenerate_count": degen_count}, f, indent=2) |
| print("wrote eval_fresh.json") |