"""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)