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