indigo-chat / chat.py
adyoi's picture
Upload chat.py with huggingface_hub
b816de0 verified
Raw
History Blame Contribute Delete
8.57 kB
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()