File size: 3,069 Bytes
8c65959
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
#!/usr/bin/env python3
"""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

# Load best.pt
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 perplexity ----
val_npy = os.path.join(OUT, "data", "val.npy")
if os.path.exists(val_npy):
    val = np.load(val_npy)
    # subsample for speed: take up to 4096 windows
    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")

# ---- generation ----
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}")

# ---- degeneracy check ----
def degenerate(text):
    # repeated n-gram loop detection
    words = text.split()
    if len(words) < 6:
        return False, "short"
    # check for 3-gram repetition covering >60% of tail
    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")