#!/usr/bin/env python """generate.py -- chat with Byrne-100M-UltraX-MC through the engine. Defaults are the PROJECT SERVING DEFAULTS: temperature sampling at temp 0.7 / top_k 40 / rep_pen 1.3. Greedy (--temp 0) is a testing tool -- see DECODING-DEFAULTS.md before using it to judge output quality. python generate.py # multi-turn REPL python generate.py -p "Explain photosynthesis." python generate.py -p "..." --device cuda --max-new 256 """ import argparse import os import sys import torch HERE = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, HERE) from spike_infer import SpikeEngine # noqa: E402 def main(): ap = argparse.ArgumentParser() ap.add_argument("-p", "--prompt", default=None) ap.add_argument("--system", default=None) ap.add_argument("--device", default="auto") ap.add_argument("--max-new", type=int, default=200) ap.add_argument("--temp", type=float, default=0.7) ap.add_argument("--top-k", type=int, default=40) ap.add_argument("--top-p", type=float, default=1.0) ap.add_argument("--rp", type=float, default=1.3) ap.add_argument("--seed", type=int, default=None, help="fixed seed for reproducibility; default is random " "per turn, as a served model behaves") ap.add_argument("--no-mtp", action="store_true", help="disable MTP speculative drafting") ap.add_argument("--raw", action="store_true", help="raw completion, no chat framing") ap.add_argument("--stats", action="store_true") ap.add_argument("--model-dir", default=None, help="alternate exported weights dir (safetensors)") ap.add_argument("--ckpt", default=None, help="released stage: checkpoints/base_62k.pt, " "checkpoints/sft_7100.pt, " "checkpoints/dpo_3200.pt (default via package.json)") args = ap.parse_args() eng = SpikeEngine(HERE, device=args.device, verbose=True, model_dir=args.model_dir, ckpt_file=args.ckpt) eng.new_conversation(system_prompt=args.system) print(f"[gen] {eng.n_params:.1f}M params ctx {eng.max_len} " f"device {eng.device} temp {args.temp} top_k {args.top_k} rp {args.rp}") def run(text, seed): fn = eng.generate_raw if args.raw else eng.chat return fn(text, max_new=args.max_new, temp=args.temp, top_k=args.top_k, top_p=args.top_p, rp=args.rp, seed=seed, use_mtp=not args.no_mtp, stream=lambda s: print(s, end="", flush=True)) def seed_for(i): if args.seed is not None: return args.seed return int.from_bytes(os.urandom(4), "little") if args.prompt is not None: out, st = run(args.prompt, seed_for(0)) print() if args.stats: print(f"[stats] {st}") return print("Multi-turn chat. Ctrl-C or empty line + 'exit' to quit.\n") i = 0 while True: try: user = input("you> ").strip() except (EOFError, KeyboardInterrupt): print() break if not user: continue if user in ("exit", "quit"): break print("bot> ", end="", flush=True) out, st = run(user, seed_for(i)) print() if args.stats: print(f"[stats] tokens={st['tokens']} {st['tok_per_s']:.1f} tok/s " f"ctx {st['context_used']}/{st['context_max']}") i += 1 if __name__ == "__main__": main()