indigo / generate.py
adyoi's picture
docs: komentar inline + README argumen + bugfix & optimasi (7dd92fc)
9a82835 verified
Raw
History Blame Contribute Delete
7.13 kB
"""
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()