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 tests/test_model.py from junafinity/Gala-598M-MLX: direct link, hf CLI and curl.
- Browser
- Download file 3.97 kB
-
https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/tests/test_model.py
- Command line
-
hf download hf://junafinity/Gala-598M-MLX/tests/test_model.py
-
curl -L -o test_model.py https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/tests/test_model.py
3.97 kB
| """Run: python -m pytest tests/ -q (or just: python tests/test_model.py)""" | |
| import sys, os | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| from mlx.utils import tree_flatten | |
| from model import (HoardConfig, HOARD, HoardMLP, chunk_gated_delta_rule, | |
| step_gated_delta_rule, inv_unit_lower, l2norm) | |
| def tiny_cfg(**kw): | |
| base = dict(vocab_size=256, d_model=64, n_cells=1, n_loops=2, n_heads=2, head_dim_k=16, | |
| head_dim_v=16, chunk_size=8, window=8, attn_heads=4, hoard_n_sub=4, | |
| hoard_block=8, hoard_topk=3, hoard_router_dim=8) | |
| base.update(kw) | |
| return HoardConfig(**base) | |
| def test_inv_unit_lower(): | |
| mx.random.seed(0) | |
| C = 16 | |
| A = mx.tril(mx.random.normal((3, C, C)), k=-1) | |
| I = mx.eye(C) | |
| prod = (I + A) @ inv_unit_lower(A) | |
| assert mx.abs(prod - I).max().item() < 1e-4 | |
| def test_chunk_vs_recurrent(): | |
| mx.random.seed(1) | |
| B, H, T, dk, dv = 2, 3, 37, 16, 12 # T not a multiple of chunk on purpose | |
| q = l2norm(mx.random.normal((B, H, T, dk))) | |
| k = l2norm(mx.random.normal((B, H, T, dk))) | |
| v = mx.random.normal((B, H, T, dv)) | |
| g = -mx.exp(mx.random.normal((B, H, T))) * 0.1 | |
| beta = mx.sigmoid(mx.random.normal((B, H, T))) | |
| o_chunk, S_chunk = chunk_gated_delta_rule(q, k, v, g, beta, chunk_size=8) | |
| S = mx.zeros((B, H, dk, dv)) | |
| outs = [] | |
| for t in range(T): | |
| o_t, S = step_gated_delta_rule(q[:, :, t], k[:, :, t], v[:, :, t], g[:, :, t], beta[:, :, t], S) | |
| outs.append(o_t) | |
| o_rec = mx.stack(outs, axis=2) | |
| assert mx.abs(o_chunk - o_rec).max().item() < 1e-4, mx.abs(o_chunk - o_rec).max().item() | |
| assert mx.abs(S_chunk - S).max().item() < 1e-4 | |
| def test_hoard_sorted_vs_unsorted(): | |
| mx.random.seed(2) | |
| cfg = tiny_cfg() | |
| m = HoardMLP(cfg) | |
| x = mx.random.normal((2, 40, cfg.d_model)) # N*k > 64 -> sorted path | |
| y_sorted = m(x) | |
| m.k_backup = m.k | |
| # force unsorted path by evaluating in small pieces | |
| ys = mx.concatenate([m(x[:, i:i + 4]) for i in range(0, 40, 4)], axis=1) | |
| assert mx.abs(y_sorted - ys).max().item() < 1e-4 | |
| def test_forward_backward_finite(): | |
| mx.random.seed(3) | |
| for mixer, mlp in (("gdn", "hoard"), ("attn", "dense"), ("gdn", "dense")): | |
| cfg = tiny_cfg(mixer=mixer, mlp=mlp) | |
| model = HOARD(cfg) | |
| ids = mx.random.randint(0, cfg.vocab_size, (2, 33)) | |
| def loss_fn(model, ids): | |
| logits = model(ids[:, :-1]) | |
| ce = nn.losses.cross_entropy(logits, ids[:, 1:], reduction="mean") | |
| return ce + model.balance_loss() | |
| loss, grads = nn.value_and_grad(model, loss_fn)(model, ids) | |
| mx.eval(loss, grads) | |
| assert mx.isfinite(loss).item() | |
| for n, g in tree_flatten(grads): | |
| assert mx.isfinite(g).all().item(), n | |
| assert loss.item() < 7.0 # ~ln(256)=5.5 + slack | |
| def test_decode_matches_prefill(): | |
| """Token-by-token decode with caches must equal one-shot logits.""" | |
| mx.random.seed(4) | |
| cfg = tiny_cfg(window=6) | |
| model = HOARD(cfg) | |
| ids = mx.random.randint(0, cfg.vocab_size, (1, 21)) | |
| full = model(ids) | |
| cache = model.new_cache() | |
| outs = [] | |
| # prefill first 5, then decode one at a time | |
| h = model.forward_hidden(ids[:, :5], cache=cache) | |
| outs.append(model.logits(h)) | |
| for t in range(5, 21): | |
| h = model.forward_hidden(ids[:, t:t + 1], cache=cache) | |
| outs.append(model.logits(h)) | |
| inc = mx.concatenate(outs, axis=1) | |
| err = mx.abs(full - inc).max().item() | |
| assert err < 1e-3, err | |
| def test_generate(): | |
| mx.random.seed(5) | |
| cfg = tiny_cfg() | |
| model = HOARD(cfg) | |
| out = model.generate(mx.random.randint(0, cfg.vocab_size, (2, 7)), max_new_tokens=9, temperature=0.8, top_k=20) | |
| assert out.shape == (2, 16) | |
| if __name__ == "__main__": | |
| for name, fn in list(globals().items()): | |
| if name.startswith("test_"): | |
| fn(); print("PASS", name) | |