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