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