min-spark / generate.py
Eclipse-Senpai's picture
scrub internal project references from generate.py
495e2c4 verified
Raw
History Blame Contribute Delete
3.03 kB
"""Raw bundled inference CLI for min-spark (no transformers dependency).
Mirrors the Space loader's generation loop exactly: EOS
prefix once, truncate to the last max_seq_len tokens, effort -> loop count.
This is the second, self-contained integration path; the Transformers path is
modeling_minspark.py. Prefer the Transformers path unless you want zero
framework overhead.
"""
from __future__ import annotations
import argparse
from pathlib import Path
import torch
from tokenizers import Tokenizer
from meiosis import Meiosis, MeiosisConfig
EFFORT_MAP = {"low": 2, "medium": 3, "high": 4}
EOS_ID = 2
HERE = Path(__file__).resolve().parent
_DEFAULT_CKPT = HERE / "model.safetensors"
_DEFAULT_TOK = HERE / "tokenizer.json"
def load_model(ckpt_path: str | None = None, device: str = "cpu") -> Meiosis:
from safetensors.torch import load_file
model = Meiosis(MeiosisConfig())
model.load_state_dict(load_file(str(ckpt_path or _DEFAULT_CKPT)), strict=False)
model.to(device).eval()
return model
@torch.no_grad()
def generate(model, tokenizer, prompt: str, *, loops: int, max_new: int,
temperature: float, top_k: int, device: str):
"""Yield decoded tokens one at a time (mirrors the Space loader)."""
ids = [EOS_ID] + tokenizer.encode(prompt).ids
for _ in range(max_new):
ctx = ids[-model.config.max_seq_len:]
x = torch.tensor([ctx], device=device)
logits = model(x, loops=loops)
next_logits = logits[0, -1] / max(temperature, 1e-6)
if top_k > 0:
topk_vals, _ = torch.topk(next_logits, min(top_k, next_logits.shape[-1]))
next_logits[next_logits < topk_vals[-1]] = float("-inf")
probs = torch.softmax(next_logits, dim=-1)
next_id = int(torch.multinomial(probs, 1).item())
if next_id == EOS_ID:
break
ids.append(next_id)
yield tokenizer.decode([next_id])
def main():
ap = argparse.ArgumentParser(description="min-spark raw inference (no transformers)")
ap.add_argument("--ckpt", default=str(_DEFAULT_CKPT))
ap.add_argument("--tokenizer", default=str(_DEFAULT_TOK))
ap.add_argument("--effort", "-e", choices=sorted(EFFORT_MAP), default="medium")
ap.add_argument("--loops", type=int, default=None)
ap.add_argument("--max-new", type=int, default=200)
ap.add_argument("--temperature", "-t", type=float, default=0.8)
ap.add_argument("--top-k", type=int, default=50)
ap.add_argument("--device", default="cpu")
ap.add_argument("--prompt", "-p", required=True)
args = ap.parse_args()
loops = args.loops if args.loops is not None else EFFORT_MAP[args.effort]
model = load_model(args.ckpt, args.device)
tok = Tokenizer.from_file(args.tokenizer)
for chunk in generate(model, tok, args.prompt, loops=loops, max_new=args.max_new,
temperature=args.temperature, top_k=args.top_k, device=args.device):
print(chunk, end="", flush=True)
print()
if __name__ == "__main__":
main()