File size: 2,831 Bytes
309aac0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fce6c24
 
 
 
 
309aac0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fce6c24
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
"""Meiosis model loader + byte-BPE tokenizer for the min-spark-preview Space.

ZeroGPU: model loads at module scope with .to("cuda") (string, never an int) so
the `spaces` hijack packs weights to disk for the GPU worker. Generation runs
inside @spaces.GPU (decorated in app.py). Vendored paths so the Space has zero
dependency on the PICO repo layout.
"""
from __future__ import annotations
from pathlib import Path
import torch

_ASSETS = Path(__file__).parent / "assets"
_TOK = _ASSETS / "tokenizer.json"
_CKPT = _ASSETS / "meiosis.safetensors"

# Keep `meiosis.py` importable: it is torch-only and self-contained here.
import sys
if str(_ASSETS) not in sys.path:
    sys.path.insert(0, str(_ASSETS))

from meiosis import Meiosis, MeiosisConfig

# tokenizer.json is a raw HF `tokenizers` artifact (no PreTrainedTokenizerFast
# wrapper, per ADR-0010) — load it with the `tokenizers` library directly.
from tokenizers import Tokenizer

EOS_ID = 2  # PICO specials: <pad>=0, <bos>=1, <eos>=2 (ADR-0010)


def load_tokenizer():
    return Tokenizer.from_file(str(_TOK))


def load_model(device: str = "cpu") -> Meiosis:
    from safetensors.torch import load_file
    model = Meiosis(MeiosisConfig())
    state = load_file(str(_CKPT))
    # strict=False: safetensors may lack non-persistent buffers (rope, loop_rms)
    model.load_state_dict(state, 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):
    """Token-by-token sampling, mirroring infer.py. Yields (chunk, count, tps)
    where `chunk` is THIS step's decoded token (not cumulative) so the caller
    can stream a typewriter effect. Per-token decode joins byte-exactly for
    this byte-level BPE (verified against cumulative decode). Runs on the GPU
    worker under @spaces.GPU; yields only CPU-safe Python objects."""
    import time
    ids = [EOS_ID] + tokenizer.encode(prompt).ids
    t0 = None
    count = 0
    for _ in range(max_new):
        ctx = ids[-model.config.max_seq_len:]
        x = torch.tensor([ctx], device=device)
        logits = model(x, loops=loops)
        if t0 is None:
            t0 = time.perf_counter()
        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)
        count += 1
        elapsed = time.perf_counter() - t0
        yield tokenizer.decode([next_id]), count, (count / elapsed if elapsed > 0 else 0.0)