Text Generation
MLX
English
apple-silicon
pretrained-from-scratch
gated-deltanet
linear-attention
product-key-memory
long-context
Instructions to use junafinity/Gala-598M-MLX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use junafinity/Gala-598M-MLX with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # if on a CUDA device, also pip install mlx[cuda] # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("junafinity/Gala-598M-MLX") prompt = "Once upon a time in" text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- MLX LM
How to use junafinity/Gala-598M-MLX with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "junafinity/Gala-598M-MLX" --prompt "Once upon a time"
- Atomic Chat
Download bench_decode.py from junafinity/Gala-598M-MLX: direct link, hf CLI and curl.
- Browser
- Download file 2.29 kB
-
https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/bench_decode.py
- Command line
-
hf download hf://junafinity/Gala-598M-MLX/bench_decode.py
-
curl -L -o bench_decode.py https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/bench_decode.py
2.29 kB
| """Decode-path comparison: tokens/sec and memory vs context length. | |
| HOARD carries constant state (GDN fast weights + a 256-token attention window) | |
| while the Transformer's KV cache grows with context. Measures batch-1 decode | |
| after prefills of increasing length. | |
| python bench_decode.py --runs runs/fw_hoard_small runs/fw_transformer_small \ | |
| --contexts 1024 4096 16384 32768 --tokens 96 | |
| """ | |
| import argparse, json, os, time | |
| import numpy as np | |
| import mlx.core as mx | |
| from model import HOARD, HoardConfig | |
| def load(run_dir): | |
| meta = json.load(open(os.path.join(run_dir, "config.json"))) | |
| cfg = HoardConfig.from_dict(meta["config"]) | |
| model = HOARD(cfg) | |
| model.load_weights(os.path.join(run_dir, "model.safetensors")) | |
| model.eval() | |
| return model, meta["args"]["config"] | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--runs", nargs="+", required=True) | |
| ap.add_argument("--contexts", type=int, nargs="+", default=[1024, 4096, 16384, 32768]) | |
| ap.add_argument("--tokens", type=int, default=96) | |
| ap.add_argument("--data", default="data/fineweb/val.bin") | |
| a = ap.parse_args() | |
| toks = np.memmap(a.data, dtype=np.uint16, mode="r") | |
| print(f"{'model':<22}{'context':>9}{'prefill s':>11}{'decode tok/s':>14}{'peak GB':>9}") | |
| for rd in a.runs: | |
| model, name = load(rd) | |
| for L in a.contexts: | |
| mx.reset_peak_memory() | |
| prompt = mx.array(np.asarray(toks[:L]).astype(np.int32))[None] | |
| cache = model.new_cache() | |
| t0 = time.time() | |
| h = model.forward_hidden(prompt, None, cache) | |
| logits = model.logits(h[:, -1:]) | |
| mx.eval(logits) | |
| t_prefill = time.time() - t0 | |
| ids = mx.argmax(logits[:, -1], axis=-1)[:, None] | |
| t0 = time.time() | |
| for _ in range(a.tokens): | |
| h = model.forward_hidden(ids, None, cache) | |
| logits = model.logits(h) | |
| ids = mx.argmax(logits[:, -1], axis=-1)[:, None] | |
| mx.eval(ids) | |
| dt = time.time() - t0 | |
| print(f"{name:<22}{L:>9}{t_prefill:>11.2f}{a.tokens / dt:>14.1f}" | |
| f"{mx.get_peak_memory() / 2**30:>9.2f}") | |
| del model | |
| mx.clear_cache() | |
| if __name__ == "__main__": | |
| main() | |