"""Interactive prompt loop for micro500. Loads the model once; exit with Ctrl+C.""" import json import sys import time import torch import torch.nn.functional as F from safetensors.torch import load_file from model import MicroLM WEIGHTS = "micro500.safetensors" CONFIG = "micro500_config.json" CONTEXT = 64 # characters of history fed to the model N_CHARS = 300 # characters generated per prompt TEMP = 0.8 TOP_K = 0 # 0 = disabled def load(): cfg = json.load(open(CONFIG)) chars = cfg["chars"] net = MicroLM(len(chars), cfg["d"], cfg["h"], cfg["pad"]) net.load_state_dict(load_file(WEIGHTS)) net.eval() stoi = {c: i for i, c in enumerate(chars)} itos = {i: c for c, i in stoi.items()} return net, stoi, itos @torch.no_grad() def generate(net, stoi, itos, prompt, n=N_CHARS, temp=TEMP, top_k=TOP_K): ids = [stoi.get(c, 0) for c in prompt.lower()] or [0] start = time.perf_counter() for _ in range(n): x = torch.tensor([ids[-CONTEXT:]]) logits = net(x)[0, -1] / temp if top_k: v, _ = torch.topk(logits, top_k) logits[logits < v[-1]] = float("-inf") ids.append(torch.multinomial(F.softmax(logits, -1), 1).item()) elapsed = time.perf_counter() - start text = "".join(itos[i] for i in ids) return text, n / elapsed def main(): net, stoi, itos = load() n_params = sum(p.numel() for p in net.parameters()) print(f"micro500 loaded ({n_params} parameters). Press Ctrl+C to quit.\n") try: while True: prompt = input("prompt> ") text, tps = generate(net, stoi, itos, prompt) print(f"\n{text}\n") print(f"[{tps:,.0f} tokens/sec]\n") except (KeyboardInterrupt, EOFError): print("\nbye") sys.exit(0) if __name__ == "__main__": main()