indigo / eval.py
adyoi's picture
docs: komentar inline + README argumen + bugfix & optimasi (7dd92fc)
9a82835 verified
Raw
History Blame Contribute Delete
7.23 kB
"""
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()