File size: 3,033 Bytes
054df1d
 
495e2c4
054df1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
495e2c4
054df1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""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()