""" 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()