palimpseste-max / examples /generate.py
thefinalboss's picture
Upload examples/generate.py with huggingface_hub
da318d2 verified
Raw
History Blame Contribute Delete
2.79 kB
#!/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()