File size: 1,804 Bytes
063093a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Standalone MLX loader and greedy/sampled generation for Sol Lassi 600K."""
from __future__ import annotations

from pathlib import Path

import mlx.core as mx
from tokenizers import Tokenizer

from keystone_mlx.dense_control import (
    DENSE_CONTROL_600K_DEEP,
    DenseControlLM,
    parameter_count,
)


def load_model(model_dir: str | Path):
    """Load the frozen 600K MLX weights and tokenizer from a local snapshot."""
    directory = Path(model_dir)
    model = DenseControlLM(DENSE_CONTROL_600K_DEEP)
    model.load_weights(str(directory / "model.npz"))
    mx.eval(model.parameters())
    if parameter_count(model) != 600_000:
        raise RuntimeError("Sol Lassi checkpoint does not contain exactly 600,000 parameters")
    tokenizer = Tokenizer.from_file(str(directory / "tokenizer.json"))
    return model, tokenizer


def generate(
    model,
    tokenizer: Tokenizer,
    prompt: str,
    max_new_tokens: int = 64,
    temperature: float = 0.0,
    seed: int = 0,
) -> str:
    """Generate a continuation; context is truncated to the trained 128 tokens."""
    if max_new_tokens < 0:
        raise ValueError("max_new_tokens must be nonnegative")
    if temperature < 0:
        raise ValueError("temperature must be nonnegative")
    mx.random.seed(seed)
    ids = tokenizer.encode(prompt).ids or [1]
    eos_id = tokenizer.token_to_id("<|eos|>")
    for _ in range(max_new_tokens):
        context = ids[-128:]
        logits = model(mx.array([context], dtype=mx.int32))[0, -1]
        if temperature > 0:
            next_id = int(mx.random.categorical(logits / temperature).item())
        else:
            next_id = int(mx.argmax(logits).item())
        ids.append(next_id)
        if eos_id is not None and next_id == eos_id:
            break
    return tokenizer.decode(ids)