#!/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()