File size: 2,945 Bytes
9e30969
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""
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