| """ |
| 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") |
|
|
| |
| parser.add_argument("--ckpt", default="out/indigo_best.safetensors", |
| help="path ke file checkpoint model (default: out/indigo_best.safetensors)") |
|
|
| |
| 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)") |
|
|
| |
| parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"], |
| help="device: auto/cpu/cuda (default: auto)") |
|
|
| |
| 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() |
|
|
| |
| 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 |
|
|
| |
| model, tokenizer = load_model(args.ckpt) |
| model = model.to(device) |
|
|
| |
| 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") |
| |
| 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)) |
| |
| 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%}") |
|
|
| |
| 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 |
|
|
| |
| if wordset is None: |
| text, _ = sample() |
| print(text) |
| return |
|
|
| |
| best_text, best_ratio = "", -1.0 |
| for t in range(args.guard_tries): |
| torch.manual_seed((args.seed or 0) + t * 1013) |
| 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 |
| print(best_text) |
| print(f"[guard] rasio kata dikenal: {best_ratio:.0%}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|