Vivid86's picture
Upload folder using huggingface_hub
9e30969 verified
Raw History Blame Contribute Delete
2.95 kB
"""
Sampler — class-based autoregressive text generation.
A thin wrapper around the model that provides a stateful `.generate()` method.
Prefer `inference/generate.py` for full feature support (top-p, EOS stopping,
KV-cache). This class is useful for:
- Interactive REPL-style generation
- Benchmarking throughput (batch generation)
- Extending with custom sampling strategies
TOP-K FIX:
The original code used `logits[logits < v.min()]` which took the global
minimum across the batch — giving some rows more than top_k candidates
and others fewer. The fix uses `v[:, -1:]` (shape [B, 1]) so each row
independently thresholds at its own k-th value.
"""
import torch
import torch.nn.functional as F
class Sampler:
"""
Stateful autoregressive token sampler.
Args:
model: A trained MiniTransformer.
tokenizer: Optional BPETokenizer for EOS stopping (pass None to disable).
device: Device to run generation on.
"""
def __init__(self, model, tokenizer=None, device: str = "cuda"):
self.model = model.to(device)
self.tokenizer = tokenizer
self.device = device
@torch.no_grad()
def generate(
self,
idx: torch.Tensor,
max_new_tokens: int = 50,
temperature: float = 1.0,
top_k: int = None,
) -> torch.Tensor:
"""
Autoregressively extend `idx` with up to `max_new_tokens` new tokens.
Args:
idx: Seed token indices, shape [B, T].
max_new_tokens: Maximum tokens to append.
temperature: Sampling temperature (0 = greedy).
top_k: Top-k filtering (None = disabled).
Returns:
Extended token tensor, shape [B, T + n_generated].
"""
self.model.eval()
context_len = self.model.config.context_len
for _ in range(max_new_tokens):
# CRITICAL: Truncate to context window — prevents pos_emb index overflow
context = idx[:, -context_len:]
logits, _ = self.model(context.to(self.device))
logits = logits[:, -1, :] / max(temperature, 1e-8) # [B, vocab]
if top_k is not None:
k = min(top_k, logits.size(-1))
v, _ = torch.topk(logits, k)
# FIX: per-row threshold via v[:, -1:] (shape [B, 1])
logits = logits.masked_fill(logits < v[:, -1:], float("-inf"))
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1) # [B, 1]
# EOS stopping (only for batch size 1)
if (
self.tokenizer is not None
and idx.size(0) == 1
and next_token.item() == self.tokenizer.eos_id
):
break
idx = torch.cat([idx, next_token], dim=1)
return idx