File size: 2,788 Bytes
da318d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
"""PALIMPSESTE — Generate text from a trained model.

Usage:
  python examples/generate.py --model ./my_model --prompt "the " --max-tokens 100
  python examples/generate.py --model ./my_model --interactive
  python examples/generate.py --model user/palimpseste-small --prompt "the "  # from HF Hub
"""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from palimseste.hf import HFPalimpsesteLM


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="Generate text from a trained PALIMPSESTE LM.")
    p.add_argument("--model", "-m", type=str, required=True,
                   help="path to a saved model dir, or an HF Hub repo id")
    p.add_argument("--prompt", "-p", type=str, default="",
                   help="prompt text (empty for BOS-only)")
    p.add_argument("--max-tokens", "-n", type=int, default=100)
    p.add_argument("--temperature", "-t", type=float, default=0.5)
    p.add_argument("--top-k", type=int, default=None)
    p.add_argument("--seed", type=int, default=None)
    p.add_argument("--interactive", "-i", action="store_true",
                   help="interactive REPL mode")
    return p.parse_args()


def load_model(model_path: str) -> HFPalimpsesteLM:
    p = Path(model_path)
    if p.exists() and p.is_dir():
        return HFPalimpsesteLM.from_pretrained(p)
    # try HF Hub
    try:
        from huggingface_hub import snapshot_download
        local = snapshot_download(repo_id=model_path, repo_type="model")
        return HFPalimpsesteLM.from_pretrained(local)
    except Exception as e:
        sys.exit(f"could not load model from '{model_path}': {e}")


def main() -> None:
    args = parse_args()
    print(f"loading model from {args.model} ...", file=sys.stderr)
    lm = load_model(args.model)
    print(f"loaded: D={lm.config.D:,}  |M|={len(lm.mem):,}  "
          f"vocab={lm.config.vocab_size}", file=sys.stderr)

    if args.interactive:
        print("PALIMPSESTE interactive mode. Ctrl+C to exit.", file=sys.stderr)
        while True:
            try:
                prompt = input(">>> ")
            except (EOFError, KeyboardInterrupt):
                print("\nbye.")
                break
            out = lm.generate(prompt, max_new_tokens=args.max_tokens,
                              temperature=args.temperature, top_k=args.top_k,
                              seed=args.seed)
            print(out.text)
    else:
        out = lm.generate(args.prompt, max_new_tokens=args.max_tokens,
                          temperature=args.temperature, top_k=args.top_k,
                          seed=args.seed)
        print(out.text)


if __name__ == "__main__":
    main()