Download sample.py from amanm10000/sprout: direct link, hf CLI and curl.
- Browser
- Download file 2.94 kB
-
https://huggingface.co/amanm10000/sprout/resolve/main/sample.py
- Command line
-
hf download hf://amanm10000/sprout/sample.py
-
curl -L -o sample.py https://huggingface.co/amanm10000/sprout/resolve/main/sample.py
2.94 kB
| """Complete a story with a locally trained Sprout checkpoint (not a chatbot).""" | |
| import argparse | |
| import contextlib | |
| import hashlib | |
| from pathlib import Path | |
| import torch | |
| from tokenizers import Tokenizer | |
| from model import Sprout | |
| def generate(model, tokenizer, prompt, max_new_tokens=180, temperature=.8, top_k=40, seed=42): | |
| if temperature <= 0 or top_k < 1 or max_new_tokens < 0: | |
| raise ValueError("temperature and top_k must be positive; token count nonnegative") | |
| device = next(model.parameters()).device | |
| generator = torch.Generator(device=device).manual_seed(seed) | |
| eot = tokenizer.token_to_id("<|endoftext|>") | |
| ids = tokenizer.encode(prompt).ids or [eot] | |
| x = torch.tensor([ids], device=device, dtype=torch.long) | |
| was_training = model.training | |
| model.eval() | |
| try: | |
| for _ in range(max_new_tokens): | |
| ctx = torch.autocast("cuda", dtype=torch.bfloat16) if device.type == "cuda" else contextlib.nullcontext() | |
| with ctx: | |
| logits = model(x[:, -model.block_size:])[:, -1, :].float() / temperature | |
| threshold = torch.topk(logits, min(top_k, logits.size(-1))).values[:, [-1]] | |
| logits = logits.masked_fill(logits < threshold, -float("inf")) | |
| token = torch.multinomial(torch.softmax(logits, dim=-1), 1, generator=generator) | |
| if token.item() == eot: | |
| break | |
| x = torch.cat((x, token), dim=1) | |
| return tokenizer.decode(x[0].tolist(), skip_special_tokens=True) | |
| finally: | |
| model.train(was_training) | |
| def main(): | |
| root = Path(__file__).resolve().parent | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--checkpoint", type=Path, default=root / "runs/sprout/best.pt") | |
| parser.add_argument("--tokenizer", type=Path, default=root / "data/tokenizer.json") | |
| parser.add_argument("--prompt", default="Once upon a time, a tiny robot found a seed.") | |
| parser.add_argument("--tokens", type=int, default=220) | |
| parser.add_argument("--temperature", type=float, default=.8) | |
| parser.add_argument("--top-k", type=int, default=40) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--device", default="cpu", choices=["cpu", "cuda"]) | |
| args = parser.parse_args() | |
| torch.set_num_threads(4) | |
| saved = torch.load(args.checkpoint, map_location="cpu", weights_only=True) | |
| expected = saved.get("tokenizer_sha256") | |
| if expected and hashlib.sha256(args.tokenizer.read_bytes()).hexdigest() != expected: | |
| raise ValueError("Tokenizer does not match checkpoint") | |
| model = Sprout(**saved["model_config"]) | |
| model.load_state_dict(saved["model"]) | |
| del saved | |
| model.to(args.device) | |
| tokenizer = Tokenizer.from_file(str(args.tokenizer)) | |
| print(generate(model, tokenizer, args.prompt, args.tokens, args.temperature, args.top_k, args.seed)) | |
| if __name__ == "__main__": | |
| main() | |