"""Interactive generator for TinyLiquid. Usage: .venv/bin/python generate.py --ckpt ckpt/forensic --persona analyst .venv/bin/python generate.py --ckpt ckpt/nlp --prompt "Once upon a time," --max-new 80 """ import argparse from pathlib import Path import torch from model.config import TinyLiquidConfig, CONFIGS from model.utils import latest_ckpt from model.tiny_liquid import TinyLiquid from data.tokenizer import load_tokenizer PERSONA_T = {"analyst": "<|analyst|>", "skeptic": "<|skeptic|>", "none": None} def parse_args(): ap = argparse.ArgumentParser() ap.add_argument("--ckpt", default="ckpt/forensic") ap.add_argument("--tok", default="data/tokenizer.json") ap.add_argument("--persona", default="analyst", choices=list(PERSONA_T)) ap.add_argument("--prompt", default=None) ap.add_argument("--max-new", type=int, default=200) ap.add_argument("--temp", type=float, default=0.8) ap.add_argument("--topk", type=int, default=40) ap.add_argument("--threads", type=int, default=8) return ap.parse_args() def main(): args = parse_args() torch.set_num_threads(args.threads) tok = load_tokenizer(args.tok) ckpt = latest_ckpt(args.ckpt) assert ckpt, f"no checkpoints in {args.ckpt}" sd = torch.load(ckpt, map_location="cpu") cfg_dict = dict(sd.get("config", CONFIGS["tiny10m"])) cfg = TinyLiquidConfig(vocab_size=tok.get_vocab_size(), **{k: v for k, v in cfg_dict.items() if k != "vocab_size"}) model = TinyLiquid(cfg) model.load_state_dict(sd["model"]) model.eval() print(f"loaded {ckpt} (step {sd.get('step','?')})", flush=True) persona_id = {"none": 0, "analyst": 1, "skeptic": 2}[args.persona] p_token = PERSONA_T[args.persona] def respond(user_text, max_new=None, temp=None): mn = max_new or args.max_new t = temp or args.temp prompt = (p_token or "") + "<|user|>" + user_text + "<|assistant|>" ids = tok.encode(prompt).ids out = model.generate(tok, ids, persona_id=persona_id, max_new=mn, temperature=t, top_k=args.topk, repetition_penalty=1.4, no_repeat_ngram_size=4) return tok.decode(out[len(ids):]) if args.prompt: print(respond(args.prompt)) return print("TinyLiquid chat. Persona:", args.persona, "| Ctrl-D to exit.") while True: try: line = input("you> ").strip() except (EOFError, KeyboardInterrupt): print() break if not line: continue print("model>", respond(line), flush=True) if __name__ == "__main__": main()