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 eval_longctx.py from junafinity/Gala-598M-MLX: direct link, hf CLI and curl.
- Browser
- Download file 4.17 kB
-
https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/eval_longctx.py
- Command line
-
hf download hf://junafinity/Gala-598M-MLX/eval_longctx.py
-
curl -L -o eval_longctx.py https://huggingface.co/junafinity/Gala-598M-MLX/resolve/main/eval_longctx.py
4.17 kB
| """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() | |