| """Text generation helpers used by scripts and loading examples.""" | |
| from __future__ import annotations | |
| import torch | |
| from .tokenizer import CharTokenizer | |
| def generate_text( | |
| model: torch.nn.Module, | |
| tokenizer: CharTokenizer, | |
| prompt: str, | |
| max_new_tokens: int = 200, | |
| temperature: float = 0.8, | |
| top_k: int | None = 40, | |
| seed: int = 1337, | |
| ) -> str: | |
| encoded = tokenizer.encode(prompt) | |
| if not encoded: | |
| raise ValueError("prompt must not be empty") | |
| input_ids = torch.tensor([encoded], dtype=torch.long, device=next(model.parameters()).device) | |
| generator = torch.Generator(device=input_ids.device).manual_seed(seed) | |
| output = model.generate( | |
| input_ids, | |
| max_new_tokens=max_new_tokens, | |
| temperature=temperature, | |
| top_k=top_k, | |
| generator=generator, | |
| ) | |
| return tokenizer.decode(output[0].tolist()) | |