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