Download inference/sampler.py from Vivid86/MiniTransformer-91M: direct link, hf CLI and curl.
- Browser
- Download file 2.95 kB
-
https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/inference/sampler.py
- Command line
-
hf download hf://Vivid86/MiniTransformer-91M/inference/sampler.py
-
curl -L -o sampler.py https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/inference/sampler.py
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 | |
| 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 | |