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
File size: 4,170 Bytes
76bbe95 | 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 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 | """Long-context evals for a checkpoint: (1) per-position CE over 8192-token
windows — if carried state helps, late positions beat early ones; (2) passkey
retrieval scored by log-prob against distractor keys at several depths.
python eval_longctx.py --run runs/c_gala_s0 --out runs/c_gala_s0/longctx.json
"""
import argparse, json, os
import numpy as np
import mlx.core as mx
import mlx.nn as nn
from model import HOARD, HoardConfig
try:
import tiktoken
ENC = tiktoken.get_encoding("gpt2")
except Exception:
ENC = None
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
def position_ce(model, toks, seq=8192, n_windows=6):
bins = [(0, 256), (256, 1024), (1024, 2048), (2048, 4096), (4096, 8192)]
sums = np.zeros(len(bins)); cnts = np.zeros(len(bins))
for w in range(n_windows):
st = w * (seq + 1)
x = np.asarray(toks[st: st + seq + 1]).astype(np.int32)[None]
if x.shape[1] < seq + 1:
break
x = mx.array(x)
logits = model(x[:, :-1])
ce = nn.losses.cross_entropy(logits.astype(mx.float32), x[:, 1:], reduction="none")[0]
mx.eval(ce)
ce = np.array(ce)
for i, (a, b) in enumerate(bins):
sums[i] += ce[a:b].sum(); cnts[i] += b - a
return {f"{a}-{b}": round(float(s / c), 4) for (a, b), s, c in zip(bins, sums, cnts)}
def seq_logprob(model, ids, tail_ids):
"""log P(tail | prefix) under the model, ids = prefix+tail."""
x = mx.array(np.array(ids, dtype=np.int32))[None]
logits = model(x[:, :-1]).astype(mx.float32)
lp = nn.log_softmax(logits[0], axis=-1)
n = len(tail_ids)
tgt = x[0, 1:]
sel = mx.take_along_axis(lp, tgt[:, None], axis=-1)[:, 0]
mx.eval(sel)
return float(np.array(sel)[-n:].sum())
def passkey(model, toks, seq=8192, depths=(0.1, 0.5, 0.9), n_trials=6, seed=0):
if ENC is None:
return {"error": "tiktoken unavailable"}
rng = np.random.default_rng(seed)
out = {}
for depth in depths:
correct = 0
for t in range(n_trials):
key = int(rng.integers(10000, 99999))
distractors = [int(rng.integers(10000, 99999)) for _ in range(4)]
needle = ENC.encode_ordinary(f" The pass key is {key}. Remember the pass key.")
query = ENC.encode_ordinary(" The pass key is")
filler_start = int(rng.integers(0, len(toks) - seq - 1))
filler = list(np.asarray(toks[filler_start: filler_start + seq]).astype(int))
pos = int(depth * (seq - len(needle) - len(query) - 24))
ctx = filler[:pos] + needle + filler[pos: seq - len(needle) - len(query) - 16]
scores = []
for cand in [key] + distractors:
tail = ENC.encode_ordinary(f" {cand}")
scores.append(seq_logprob(model, ctx + query + tail, tail))
correct += int(np.argmax(scores) == 0)
out[f"depth_{depth}"] = round(correct / n_trials, 3)
return out
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--run", required=True)
ap.add_argument("--data", default="data/fineweb/val.bin")
ap.add_argument("--out", default=None)
ap.add_argument("--trials", type=int, default=6)
ap.add_argument("--depths", type=float, nargs="+", default=[0.1, 0.5, 0.9])
ap.add_argument("--seqs", type=int, nargs="+", default=[8192, 32768],
help="context lengths for the passkey test")
a = ap.parse_args()
model = load(a.run)
toks = np.memmap(a.data, dtype=np.uint16, mode="r")
res = {"position_ce": position_ce(model, toks),
"passkey_acc": {str(sq): passkey(model, toks, seq=sq, depths=tuple(a.depths),
n_trials=a.trials)
for sq in a.seqs}}
out = a.out or os.path.join(a.run, "longctx.json")
json.dump(res, open(out, "w"), indent=2)
print(json.dumps(res))
if __name__ == "__main__":
main()
|