Sol-Lassi / modeling_sol_lassi.py
j0no12's picture
Publish Sol Lassi 600K Base
063093a verified
Raw History Blame Contribute Delete
1.8 kB
"""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)