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