""" chat_sft.py -- interactive multi-turn chat REPL for an SFT'd Quazimoto-LM. Renders the running conversation in the student's ChatML and generates the assistant reply. Two decode paths: * SPECULATIVE (default) -- DeepSpec/DSpark-style self-speculative decoding: the MTP heads draft the next few tokens, the main head verifies them in one pass, accepting the longest correct prefix (see model.forward_drafts / generate.py generate_speculative). Greedy, so output is deterministic. * SAMPLED (--temperature > 0) -- the KV-cache sampler (generate.py generate): the conversation is prefilled once, then each new token is a single-position cached forward. Supports temperature / top-k / top-p. Both reuse the model's KV cache; speculative additionally drafts multiple tokens per verify. Usage: python chat_sft.py # auto-find newest chkpt, speculative python chat_sft.py --temperature 0.7 # sampled instead python chat_sft.py --ckpt chkpt/quazimoto_sft.pt --system "You are concise." """ import argparse, os, sys # Force UTF-8 console so byte-merge tokens with non-cp1252 chars don't crash print(). for _s in (sys.stdout, sys.stderr): try: _s.reconfigure(encoding="utf-8", errors="replace") except (AttributeError, ValueError): pass import torch from model import QuazimotoLM, QuazimotoConfig from generate import generate, generate_speculative, resolve_stop_ids PKG_DIR = os.path.dirname(os.path.abspath(__file__)) ROLE_TOK = {"system": "<|system|>", "user": "<|user|>", "assistant": "<|assistant|>"} def find_ckpt(path): if path and os.path.isfile(path): return path folder = (os.path.dirname(path) if path else os.path.join(PKG_DIR, "chkpt")) or "." if not os.path.isdir(folder): return None # prefer an SFT checkpoint, else newest *.pt pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")] sft = [p for p in pts if "sft" in os.path.basename(p).lower()] pool = sft or pts return max(pool, key=os.path.getmtime) if pool else None def load_tokenizer(tok_dir): sys.path.insert(0, tok_dir) from spike_tokenizer import SpikeTokenizer return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json")) def render(history, tok): """history: list of {role,content} -> token ids ending with the assistant header so the model continues the assistant turn.""" parts = [f"<|im_start|>{ROLE_TOK.get(m['role'],'<|user|>')}\n{m['content']}<|im_end|>\n" for m in history] parts.append("<|im_start|><|assistant|>\n") return tok.encode("".join(parts), add_special_tokens=False) def main(): p = argparse.ArgumentParser(description="Quazimoto-LM SFT chat (KV cache + speculative)") p.add_argument("--ckpt", default="", help="SFT checkpoint (default: newest *sft*.pt in ./chkpt)") p.add_argument("--tok_dir", default=PKG_DIR) p.add_argument("--system", default="", help="optional system prompt") p.add_argument("--max_new_tokens", type=int, default=256) p.add_argument("--temperature", type=float, default=0.0, help="0 = speculative/greedy; >0 = sampled") p.add_argument("--top_k", type=int, default=40) p.add_argument("--top_p", type=float, default=0.95) p.add_argument("--repetition_penalty", type=float, default=1.1) p.add_argument("--no_speculative", action="store_true", help="force the cached sampler even at temp 0") p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") args = p.parse_args() path = find_ckpt(args.ckpt) if path is None: print("No checkpoint found; pass --ckpt or train one."); return ckpt = torch.load(path, map_location=args.device, weights_only=False) cfg = QuazimotoConfig(**ckpt["family_config"]) model = QuazimotoLM(cfg); model.load_state_dict(ckpt["model"], strict=False) model.to(args.device).eval() tok = load_tokenizer(args.tok_dir) stop = resolve_stop_ids(tok, chat=True) # <|im_end|>, , <|endoftext|> imend = tok.get_vocab().get("<|im_end|>") spec = (args.temperature <= 1e-4 and not args.no_speculative and model.mtp_heads is not None) mode = "speculative (DSpark draft+verify)" if spec else f"sampled (temp {args.temperature}, KV cache)" print(f"loaded {os.path.basename(path)} (step {ckpt.get('step')}) | decode: {mode}") print("type your message; /reset clears history, /quit exits.\n") history = [] if args.system: history.append({"role": "system", "content": args.system}) while True: try: user = input("you> ").strip() except (EOFError, KeyboardInterrupt): print(); break if not user: continue if user in ("/quit", "/exit"): break if user == "/reset": history = ([{"role": "system", "content": args.system}] if args.system else []) print("(history cleared)\n"); continue history.append({"role": "user", "content": user}) ids = render(history, tok) x = torch.tensor([ids], device=args.device) with torch.no_grad(): if spec: out = generate_speculative(model, cfg, x, args.max_new_tokens, stop_ids=stop) else: out = generate(model, cfg, x, args.max_new_tokens, temperature=max(args.temperature, 1e-6), top_k=args.top_k or None, top_p=args.top_p, repetition_penalty=args.repetition_penalty, stop_ids=stop, device=args.device, use_cache=True) gen = out[0, len(ids):].tolist() if imend in gen: # trim at the assistant turn end gen = gen[:gen.index(imend)] reply = tok.decode(gen, skip_special_tokens=True).strip() try: print(f"bot> {reply}\n") except UnicodeEncodeError: print("bot> " + reply.encode("ascii", "replace").decode() + "\n") history.append({"role": "assistant", "content": reply}) if __name__ == "__main__": main()