File size: 7,130 Bytes
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f2265aa
6326055
c7c5a39
6326055
f2265aa
 
c7c5a39
6326055
 
 
045f351
9a82835
 
 
 
 
 
 
 
 
 
 
 
045f351
 
 
 
c7c5a39
 
 
 
 
 
 
 
 
 
 
 
 
 
045f351
 
6326055
 
9a82835
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6326055
 
9a82835
6326055
 
045f351
6326055
9a82835
c7c5a39
045f351
6326055
9a82835
4551faf
3201cb8
4551faf
3201cb8
 
4551faf
 
 
3201cb8
9a82835
3201cb8
 
 
 
9a82835
3201cb8
 
 
 
 
 
4551faf
9a82835
6326055
045f351
4551faf
 
9a82835
 
 
 
 
4551faf
 
 
 
 
 
 
 
 
3201cb8
4551faf
 
9a82835
4551faf
 
 
 
 
9a82835
4551faf
 
9a82835
4551faf
 
 
 
 
9a82835
4551faf
 
6326055
 
 
 
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
"""
Script inferensi / generasi teks dari checkpoint Indigo.

Mendukung:
- Top-k, top-p, temperature, repetition penalty sampling
- KV-cache untuk generasi cepat (token-by-token)
- Guard kamus: generate beberapa kandidat β†’ pilih yang rasio kata dikenal tertinggi
- Guard morfologi: cek imbuhan Indonesia (prefiks + sufiks + asimilasi)

Cara pakai:
    python generate.py --prompt "Indigo" --max-new 300
    python generate.py --prompt "hello" --temperature 0.8 --top-k 40 --top-p 0.9
    python generate.py --prompt "kepekaan" --guard data/kamus_id.txt --guard-min 0.6
"""

import sys
import torch
import argparse

sys.stdout.reconfigure(encoding="utf-8", errors="replace")

from indigo.common import load_meta, build_tokenizer
from indigo.model import GPT, GPTConfig


def load_model(path):
    """Muat model GPT + tokenizer dari file checkpoint.

    Mendukung dua format:
    1. .safetensors: format utama Indigo
    2. .pt: format PyTorch lama

    Args:
        path: Path ke file checkpoint.

    Returns:
        Tuple (model, tokenizer).
    """
    if path.endswith(".safetensors"):
        from safetensors.torch import load_file

        state = load_file(path)
        meta = load_meta(path)
        config_d = meta["config"]
        tinfo = meta.get("tokenizer") or {"type": "char"}
        vocab = meta.get("vocab")
    else:
        ckpt = torch.load(path, map_location="cpu", weights_only=True)
        state = ckpt["model"]
        config_d = ckpt["config"]
        tinfo = ckpt.get("tokenizer") or {"type": "char"}
        vocab = ckpt.get("vocab")
    tokenizer = build_tokenizer(tinfo, vocab)
    model = GPT(GPTConfig(**config_d))
    model.load_state_dict(state, strict=False)
    return model, tokenizer


def main():
    parser = argparse.ArgumentParser(description="Generate teks dari checkpoint Indigo")

    # --- Checkpoint ---
    parser.add_argument("--ckpt", default="out/indigo_best.safetensors",
                        help="path ke file checkpoint model (default: out/indigo_best.safetensors)")

    # --- Prompt & Generasi ---
    parser.add_argument("--prompt", default="",
                        help="teks awal (prompt) untuk memulai generasi (default: kosong)")
    parser.add_argument("--max-new", type=int, default=300,
                        help="jumlah token baru yang akan dihasilkan (default: 300)")
    parser.add_argument("--temperature", type=float, default=0.8,
                        help="skala randomness: 0.0 β‰ˆ greedy, 0.8 β‰ˆ standar, >1.0 β‰ˆ random (default: 0.8)")
    parser.add_argument("--top-k", type=int, default=40,
                        help="batasi sampling ke k token teratas (0 = nonaktif, default: 40)")
    parser.add_argument("--top-p", type=float, default=1.0,
                        help="nucleus sampling: batasi kumulatif probabilitas (1.0 = nonaktif, default: 1.0)")
    parser.add_argument("--repetition-penalty", type=float, default=1.0,
                        help="penalti pengulangan token (>1.0 = aktif, 1.0 = nonaktif, default: 1.0)")
    parser.add_argument("--seed", type=int, default=None,
                        help="seed random (None = tidak ditentukan, default: None)")

    # --- Device ---
    parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"],
                        help="device: auto/cpu/cuda (default: auto)")

    # --- Guard Kamus ---
    parser.add_argument("--guard", default=None,
                        help="path file kamus (satu kata per baris); generate beberapa kandidat β†’ pilih terbaik")
    parser.add_argument("--guard-prefiks", default=None,
                        help="path file prefiks Indonesia (default: data/prefiks.txt bila ada)")
    parser.add_argument("--guard-sufiks", default=None,
                        help="path file sufiks Indonesia (default: data/sufiks.txt bila ada)")
    parser.add_argument("--guard-tries", type=int, default=5,
                        help="jumlah kandidat generate saat --guard aktif (default: 5)")
    parser.add_argument("--guard-min", type=float, default=0.6,
                        help="rasio kata dikenal minimum β€” berhenti generate jika tercapai (default: 0.6)")

    args = parser.parse_args()

    # --- Setup seed & device ---
    if args.seed is not None:
        torch.manual_seed(args.seed)
    device = "cuda" if torch.cuda.is_available() else "cpu" if args.device == "auto" else args.device

    # --- Muat model ---
    model, tokenizer = load_model(args.ckpt)
    model = model.to(device)

    # --- Muat kamus (jika --guard aktif) ---
    wordset = None
    pref_set = suf_set = None
    if args.guard:
        from pathlib import Path as _Path

        from indigo.common import load_wordlist, word_known_ratio

        wordset = load_wordlist(args.guard)
        p_def, s_def = _Path("data/prefiks.txt"), _Path("data/sufiks.txt")
        # Muat prefiks: prioritaskan argumen CLI β†’ default path
        if args.guard_prefiks and _Path(args.guard_prefiks).exists():
            pref_set = load_wordlist(args.guard_prefiks)
        elif not args.guard_prefiks and p_def.exists():
            pref_set = load_wordlist(str(p_def))
        # Muat sufiks: prioritaskan argumen CLI β†’ default path
        if args.guard_sufiks and _Path(args.guard_sufiks).exists():
            suf_set = load_wordlist(args.guard_sufiks)
        elif not args.guard_sufiks and s_def.exists():
            suf_set = load_wordlist(str(s_def))
        mode = "dengan formula afiks" if pref_set and suf_set else "kata persis"
        print(f"[guard] kamus: {len(wordset):,} kata ({mode}) | target rasio >= {args.guard_min:.0%}")

    # --- Encode prompt β†’ token IDs ---
    ids = tokenizer.encode(args.prompt) or [0]
    idx = torch.tensor([ids], dtype=torch.long, device=device)

    def sample():
        """Generate satu kandidat teks dari model.

        Returns:
            Tuple (text, ratio) β€” teks hasil generate dan rasio kata dikenal.
        """
        out = model.generate(
            idx,
            args.max_new,
            temperature=args.temperature,
            top_k=args.top_k,
            top_p=args.top_p,
            repetition_penalty=args.repetition_penalty,
        )
        text = tokenizer.decode(out[0].tolist())
        ratio = word_known_ratio(text, wordset, pref_set, suf_set) if wordset else 1.0
        return text, ratio

    # --- Tanpa guard: langsung generate & print ---
    if wordset is None:
        text, _ = sample()
        print(text)
        return

    # --- Dengan guard: generate beberapa kandidat β†’ pilih yang terbaik ---
    best_text, best_ratio = "", -1.0
    for t in range(args.guard_tries):
        torch.manual_seed((args.seed or 0) + t * 1013)  # seed berbeda tiap kandidat
        text, ratio = sample()
        mark = f"  [kandidat {t + 1}: {ratio:.0%}]"
        if ratio > best_ratio:
            best_text, best_ratio = text, ratio
        if best_ratio >= args.guard_min:
            break  # sudah cukup bagus, tidak perlu generate lagi
    print(best_text)
    print(f"[guard] rasio kata dikenal: {best_ratio:.0%}")


if __name__ == "__main__":
    main()