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