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()