File size: 7,234 Bytes
9a82835 b05971d 9a82835 b05971d 3ff4730 9a82835 3ff4730 b05971d 9a82835 3ff4730 9a82835 3ff4730 9a82835 3ff4730 9a82835 3ff4730 9a82835 b05971d 3ff4730 9a82835 3ff4730 9a82835 b05971d 9a82835 b05971d 9a82835 b05971d 9a82835 b05971d 9a82835 b05971d 9a82835 b05971d 9a82835 b05971d 3ff4730 9a82835 b05971d 9a82835 b05971d 538b95e b05971d 9a82835 b05971d 9a82835 3ff4730 b05971d | 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 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | """
Script evaluasi batched untuk membandingkan checkpoint Indigo.
Menghitung metrik per-token (nats/token) dan per-karakter (nats/karakter) pada
set uji tetap, sehingga run dengan tokenizer berbeda (char vs BPE) tetap sebanding.
Alur kerja:
1. Muat teks uji (data/sample.txt atau custom)
2. Untuk setiap checkpoint:
a. Muat model + tokenizer
b. Encode teks uji β pecah menjadi jendela (block_size)
c. Hitung cross-entropy loss per-token secara batched
d. Konversi ke nats/karakter menggunakan compression ratio
3. (Opsional) Generate teks β hitung rasio kata dikenal (kamus guard)
Cara pakai:
python eval.py --ckpt out/indigo_best.safetensors
python eval.py --ckpt runs/*/ckpt/indigo.safetensors --test data/sample.txt
python eval.py --ckpt out/indigo_best.safetensors --guard data/kamus_id.txt
"""
import argparse
from pathlib import Path
import torch
from indigo.common import (
build_tokenizer,
load_meta,
load_wordlist,
read_clean,
word_known_ratio,
)
from indigo.model import GPT, GPTConfig
def muat(path):
"""Muat model GPT + tokenizer dari checkpoint .safetensors.
Args:
path: Path ke file .safetensors.
Returns:
Tuple (model, tokenizer, meta).
"""
from safetensors.torch import load_file
meta = load_meta(path)
model = GPT(GPTConfig(**meta["config"]))
missing, unexpected = model.load_state_dict(load_file(path), strict=False)
if missing or unexpected:
raise SystemExit(f"bobot tidak cocok untuk {path}: {missing[:3]} {unexpected[:3]}")
model.eval()
tokenizer = build_tokenizer(meta.get("tokenizer") or {"type": "char"}, meta.get("vocab"))
return model, tokenizer, meta
@torch.no_grad()
def nats_per_token(model, ids, block_size, device, batch_size=32):
"""Hitung loss rata-rata (nats per token) pada seluruh sequence.
Algoritma:
1. Pecah sequence panjang menjadi jendela-jendela sepanjang block_size
2. Pad jendela ke panjang yang sama dalam batch (zero-padding + mask)
3. Forward pass batched β hitung cross-entropy per token β rata-rata
Args:
model: Model GPT.
ids: List of int β token IDs dari teks uji.
block_size: Int β panjang konteks model.
device: Str β "cpu" atau "cuda".
batch_size: Int β jumlah jendela per batch (default: 32).
Returns:
Tuple (nats_per_token, total_tokens).
"""
# Pecah sequence menjadi jendela-jendela block_size
jendela = []
for i in range(0, max(0, len(ids) - 1), block_size):
potongan = ids[i : i + block_size + 1] # +1 untuk target
if len(potongan) >= 2:
jendela.append(potongan)
# Proses batched
total_nll = 0.0
total_tok = 0
for k in range(0, len(jendela), batch_size):
kelompok = jendela[k : k + batch_size]
L = max(len(w) - 1 for w in kelompok) # panjang terpanjang dalam batch
# Buat tensor x (input), y (target), dan mask (ignore padding)
x = torch.zeros(len(kelompok), L, dtype=torch.long)
y = torch.zeros(len(kelompok), L, dtype=torch.long)
mask = torch.zeros(len(kelompok), L, dtype=torch.bool)
for r, w in enumerate(kelompok):
n = len(w) - 1
x[r, :n] = torch.tensor(w[:-1], dtype=torch.long) # input: semua kecuali terakhir
y[r, :n] = torch.tensor(w[1:], dtype=torch.long) # target: semua kecuali pertama
mask[r, :n] = True # hanya hitung posisi yang ada isinya
x, y, mask = x.to(device), y.to(device), mask.to(device)
# Forward pass β log probability β negative log-likelihood
logits, _ = model(x)
logp = torch.log_softmax(logits.float(), dim=-1)
nll = -logp.gather(2, y.unsqueeze(2)).squeeze(2)
# Akumulasi (hanya hitung posisi yang dimask)
total_nll += float(nll[mask].sum())
total_tok += int(mask.sum())
return total_nll / max(1, total_tok), total_tok
def main():
ap = argparse.ArgumentParser(
description="Skor checkpoint pada set uji tetap agar antar-run dapat dibandingkan"
)
ap.add_argument("--ckpt", nargs="+", required=True,
help="path ke satu atau lebih file checkpoint (.safetensors)")
ap.add_argument("--test", default="data/sample.txt",
help="path ke file teks uji (default: data/sample.txt)")
ap.add_argument("--device", default="cpu", choices=["cpu", "cuda"],
help="device untuk evaluasi (default: cpu)")
ap.add_argument("--guard", default=None,
help="path ke file kamus; generate teks β hitung rasio kata dikenal")
ap.add_argument("--guard-max-new", type=int, default=120,
help="jumlah token generate untuk evaluasi guard (default: 120)")
ap.add_argument("--seed", type=int, default=42,
help="seed untuk generate saat --guard aktif (default: 42)")
ap.add_argument("--batch-size", type=int, default=32,
help="batch size untuk evaluasi (default: 32)")
args = ap.parse_args()
# --- Muat teks uji ---
teks = read_clean(args.test)
n_karakter = len(teks.encode("utf-8"))
# --- Muat kamus (jika --guard) ---
wordset = prefiks = sufiks = None
if args.guard:
root = Path(__file__).resolve().parent
wordset = load_wordlist(args.guard)
p, s = root / "data" / "prefiks.txt", root / "data" / "sufiks.txt"
prefiks = load_wordlist(str(p)) if p.exists() else None
sufiks = load_wordlist(str(s)) if s.exists() else None
# --- Header tabel ---
print(f"set uji: {args.test} ({n_karakter:,} karakter)")
print(f"{'checkpoint':44s} {'nats/tok':>9s} {'nat/kar':>8s} {'kamus':>7s}")
baris = []
# --- Evaluasi setiap checkpoint ---
for path in args.ckpt:
model, tokenizer, meta = muat(path)
# Encode teks uji β hitung nats per token
ids = tokenizer.encode(teks)
npt, n_tok = nats_per_token(
model, ids, meta["config"]["block_size"], args.device, args.batch_size
)
# Konversi: nats/token β nats/karakter (menggunakan compression ratio)
kompresi = n_karakter / max(1, n_tok)
npc = npt / kompresi
# (Opsional) hitung rasio kata dikenal via generate
rasio = ""
if wordset:
torch.manual_seed(args.seed)
out = model.generate(
torch.tensor([[0]], dtype=torch.long, device=args.device),
args.guard_max_new,
temperature=0.8,
top_k=40,
)
teks_out = tokenizer.decode(out[0].tolist())
rasio = f"{word_known_ratio(teks_out, wordset, prefiks, sufiks):6.0%}"
# Format nama checkpoint yang pendek (runs/xxx/ckpt/file.safetensors)
bagian = str(Path(path)).replace("\\", "/").split("/")
nama = "/".join(bagian[-3:-1] + [bagian[-1]]) if len(bagian) >= 3 else bagian[-1]
print(f"{nama:44s} {npt:9.3f} {npc:8.3f} {rasio:>7s}")
baris.append({"ckpt": str(path), "nats_per_token": round(npt, 4), "nats_per_char": round(npc, 4)})
if __name__ == "__main__":
main()
|