OxMini / src /oxmini /generate.py
Shivam3002's picture
Publish trained OxMini checkpoint and measured model card
46144df verified
Raw
History Blame Contribute Delete
882 Bytes
"""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())