| import argparse |
| import json |
| import os |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| ROOT = Path(__file__).resolve().parent |
| INDIGO_TORCH = ROOT.parent / "Indigo" |
| INDIGO_TF = ROOT.parent / "indigo.tf" |
|
|
|
|
| def meta_of(ckpt): |
| meta_path = str(ckpt)[: -len(".safetensors")] + "_meta.json" |
| if not os.path.exists(meta_path): |
| raise SystemExit(f"meta tidak ditemukan: {meta_path}") |
| with open(meta_path, encoding="utf-8") as f: |
| return json.load(f) |
|
|
|
|
| def load_torch(ckpt, meta, device): |
| sys.path.insert(0, str(INDIGO_TORCH)) |
| from safetensors.torch import load_file |
|
|
| from indigo.common import build_tokenizer |
| from indigo.model import GPT, GPTConfig |
|
|
| state = load_file(str(ckpt)) |
| model = GPT(GPTConfig(**meta["config"])) |
| missing, unexpected = model.load_state_dict(state, strict=False) |
| if missing or unexpected: |
| print(f"state_dict: missing={missing} unexpected={unexpected}") |
| model = model.to(device) |
| tokenizer = build_tokenizer(meta.get("tokenizer") or {"type": "char"}, meta.get("vocab")) |
| return model, tokenizer, "torch" |
|
|
|
|
| def load_tf(ckpt, meta): |
| sys.path.insert(0, str(INDIGO_TF)) |
| import numpy as np |
| import tensorflow as tf |
| from safetensors.numpy import load_file |
|
|
| from indigotf.common import build_tokenizer |
| from indigotf.model import build_gpt, generate as tf_generate |
|
|
| keys = ("vocab_size", "block_size", "n_layer", "n_head", "n_embd", "dropout") |
| model = build_gpt(**{k: meta["config"][k] for k in keys}) |
| state = {k.replace("/", "_"): v for k, v in load_file(str(ckpt)).items()} |
| by_path = {v.path.replace("/", "_"): v.path for v in model.weights} |
| missing = [p for p in by_path if p not in state] |
| if missing: |
| raise SystemExit(f"bobot tidak cocok: {missing[:5]}") |
| model.set_weights([state[v.path.replace("/", "_")] for v in model.weights]) |
| tokenizer = build_tokenizer(meta.get("tokenizer") or {"type": "char"}, meta.get("vocab")) |
| n_params = int(sum(int(np.prod(v.shape)) for v in model.weights)) |
| print(f"(tensorflow dimuat, params={n_params / 1e6:.2f}M)") |
| return model, tokenizer, "tf", tf_generate |
|
|
|
|
| def mean_nll(logp_fn, context, new_ids, block_size): |
| if not new_ids: |
| return float("inf") |
| seq = (context + list(new_ids))[-block_size:] |
| T = len(seq) |
| w = min(len(new_ids), max(T - 1, 1)) |
| targets = np.asarray(seq[-w:]) |
| logp = logp_fn(seq) |
| rows = np.arange(T - 1 - w, T - 1) |
| return float(-logp[rows, targets].mean()) |
|
|
|
|
| def main(): |
| global INDIGO_TORCH, INDIGO_TF |
| sys.stdout.reconfigure(encoding="utf-8", errors="replace") |
| parser = argparse.ArgumentParser(description="Chat tester untuk model Indigo (PyTorch & TensorFlow)") |
| parser.add_argument("--ckpt", default=None, help="path .safetensors (deteksi backend otomatis)") |
| parser.add_argument("--torch-dir", default=str(INDIGO_TORCH)) |
| parser.add_argument("--tf-dir", default=str(INDIGO_TF)) |
| parser.add_argument("--max-new", type=int, default=120) |
| parser.add_argument("--temperature", type=float, default=0.8) |
| parser.add_argument("--top-k", type=int, default=40) |
| parser.add_argument("--top-p", type=float, default=0.95) |
| parser.add_argument("--repetition-penalty", type=float, default=1.15) |
| parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"]) |
| parser.add_argument("--fallback", default="saya tidak punya data.", |
| help='jawaban saat model tidak yakin; "" untuk mematikan') |
| parser.add_argument("--threshold", type=float, default=None, |
| help="ambang NLL prompt per token (default: val_loss - 0.6)") |
| parser.add_argument("--guard", default=None, help="kamus kata (satu/baris); fallback jika ratio rendah") |
| parser.add_argument("--guard-min", type=float, default=0.5, help="ambang rasio kata dikenal") |
| parser.add_argument("--show-nll", action="store_true", help="tampilkan skor NLL tiap jawaban") |
| args = parser.parse_args() |
|
|
| INDIGO_TORCH = Path(args.torch_dir) |
| INDIGO_TF = Path(args.tf_dir) |
|
|
| ckpt = args.ckpt |
| if ckpt is None: |
| cand = INDIGO_TORCH / "out" / "indigo_best.safetensors" |
| if cand.exists(): |
| ckpt = cand |
| else: |
| raise SystemExit("tidak ada checkpoint; gunakan --ckpt") |
| ckpt = Path(ckpt) |
| meta = meta_of(ckpt) |
| backend = meta.get("backend", "pytorch") |
|
|
| wordset = prefiks = sufiks = None |
| if args.guard: |
| sys.path.insert(0, str(INDIGO_TORCH)) |
| from indigo.common import load_wordlist as _load |
|
|
| wordset = _load(args.guard) |
| p_def = INDIGO_TORCH / "data" / "prefiks.txt" |
| s_def = INDIGO_TORCH / "data" / "sufiks.txt" |
| prefiks = _load(str(p_def)) if p_def.exists() else None |
| sufiks = _load(str(s_def)) if s_def.exists() else None |
|
|
| if backend == "tensorflow": |
| model, tokenizer, backend_label, tf_generate_fn = load_tf(ckpt, meta) |
| import tensorflow as tf |
|
|
| def idx_input(ctx): |
| return tf.constant([ctx], dtype=tf.int64) |
|
|
| def logp_fn(window): |
| logits = model(tf.constant([window], dtype=tf.int64), training=False) |
| return tf.nn.log_softmax(tf.cast(logits[0], tf.float32), axis=-1).numpy() |
|
|
| def generate_fn(mdl, idx, max_new, block_size_unused, temperature=1.0, top_k=None): |
| return tf_generate_fn(mdl, idx, max_new, meta["config"]["block_size"], |
| temperature=temperature, top_k=top_k) |
| else: |
| import torch |
|
|
| device = ( |
| ("cuda" if torch.cuda.is_available() else "cpu") if args.device == "auto" else args.device |
| ) |
| model, tokenizer, backend_label = load_torch(ckpt, meta, device) |
| dev = next(model.parameters()).device |
|
|
| def idx_input(ctx): |
| return torch.tensor([ctx], dtype=torch.long, device=dev) |
|
|
| def logp_fn(window): |
| idx = torch.tensor([window], dtype=torch.long, device=dev) |
| with torch.no_grad(): |
| logits, _ = model(idx) |
| return torch.log_softmax(logits[0].float(), dim=-1).cpu().numpy() |
|
|
| def generate_fn(mdl, idx, max_new, block_size_unused, temperature=1.0, top_k=None): |
| return mdl.generate(idx, max_new, temperature=temperature, top_k=top_k, |
| top_p=args.top_p, repetition_penalty=args.repetition_penalty) |
|
|
| if args.threshold is None: |
| val = meta.get("val_loss") |
| args.threshold = max(3.0, val - 0.6) if val else 5.0 |
| info_lines = [ |
| f"model={ckpt.name} | backend={backend_label} | " |
| f"tokenizer={meta.get('tokenizer', {}).get('type', 'char')} | ambang nll={args.threshold:.2f}", |
| "perintah: /reset ulang konteks | /keluar berhenti", |
| ] |
| if wordset: |
| info_lines.append(f"[guard] kamus: {len(wordset):,} kata | ambang ratio >= {args.guard_min:.0%}") |
| print("\n".join(info_lines) + "\n") |
|
|
| block_size = meta["config"]["block_size"] |
| history = [] |
|
|
| def respond(user_text): |
| piece = tokenizer.encode("\nAnda: " + user_text + "\nIndigo:") |
| context = (history + piece)[-block_size:] |
| prev_history = list(history) |
| prompt_nll = mean_nll(logp_fn, [], piece, block_size) |
| out = generate_fn(model, idx_input(context), args.max_new, block_size, |
| temperature=args.temperature, top_k=args.top_k) |
| full = (out[0].tolist() if hasattr(out[0], "tolist") else out[0]) |
| new_ids = full[len(context):] |
| text = tokenizer.decode(new_ids).strip() |
| ratio = word_known_ratio(text, wordset, prefiks, sufiks) if wordset else 1.0 |
| guard_ok = ratio >= args.guard_min if wordset else True |
| if args.fallback and (not text or prompt_nll > args.threshold or not guard_ok): |
| history[:] = prev_history |
| return args.fallback, prompt_nll |
| history.clear() |
| history.extend(full[-block_size:]) |
| return text, prompt_nll |
|
|
| while True: |
| try: |
| user = input("\nAnda> ").strip() |
| except (EOFError, KeyboardInterrupt): |
| print() |
| break |
| if not user: |
| continue |
| if user in ("/keluar", "/quit", "/exit"): |
| break |
| if user == "/reset": |
| history.clear() |
| print("(konteks direset)") |
| continue |
| reply, score = respond(user) |
| suffix = f" [nll={score:.2f}]" if args.show_nll else "" |
| print(f"Indigo> {reply}{suffix}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|