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: 1,295 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 | """
Generate text from a checkpoint with constant-memory decode.
python sample.py --run runs/hoard_small_XXXX --prompt "The dragon" --tokens 100
"""
import argparse, json, os
import mlx.core as mx
from model import HOARD, HoardConfig
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--run", required=True)
ap.add_argument("--prompt", default="\n")
ap.add_argument("--tokens", type=int, default=100)
ap.add_argument("--temperature", type=float, default=0.8)
ap.add_argument("--top_k", type=int, default=50)
ap.add_argument("--loops", type=int, default=None, help="try fewer/more loops than trained")
ap.add_argument("--seed", type=int, default=0)
a = ap.parse_args()
meta = json.load(open(os.path.join(a.run, "config.json")))
cfg = HoardConfig.from_dict(meta["config"])
model = HOARD(cfg)
model.load_weights(os.path.join(a.run, "model.safetensors"))
model.eval()
mx.random.seed(a.seed)
import tiktoken
enc = tiktoken.get_encoding("gpt2")
ids = mx.array([enc.encode_ordinary(a.prompt)])
out = model.generate(ids, max_new_tokens=a.tokens, temperature=a.temperature,
top_k=a.top_k, n_loops=a.loops)
print(enc.decode(out[0].tolist()))
if __name__ == "__main__":
main()
|