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