File size: 6,043 Bytes
19ca28d | 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 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | # model/sampling.py
#
# Token sampling for PyCraft-1 generation.
#
# Kept in its own module so the cached and uncached decode paths provably
# share one RNG consumption pattern — that is what makes the cached-vs-
# uncached equivalence test meaningful. If the two paths drew randomness
# differently, an equivalence failure could not distinguish a cache bug from
# a sampling bug.
#
# Filter order is deliberate and matches HuggingFace:
#
# repetition penalty (on raw logits)
# -> temperature
# -> top-k
# -> top-p
# -> softmax
# -> multinomial
#
# Applying the penalty after temperature would change its effective strength;
# applying top-p before top-k would change the candidate set it sees.
import torch
# ------------------------------------------------------------------ #
# Individual filters
# ------------------------------------------------------------------ #
def apply_repetition_penalty(
logits: torch.Tensor, # (batch, vocab)
prev_ids: torch.Tensor, # (batch, n_prev)
penalty: float,
) -> torch.Tensor:
"""
Discourage tokens that have already appeared (CTRL / HuggingFace formula).
Positive logits are divided by the penalty and negative ones multiplied,
so both move toward -inf regardless of sign.
"""
if penalty == 1.0:
return logits
score = torch.gather(logits, 1, prev_ids)
score = torch.where(score < 0, score * penalty, score / penalty)
return logits.scatter(1, prev_ids, score)
def top_k_filter(logits: torch.Tensor, k: int | None) -> torch.Tensor:
"""Keep only the k highest-scoring tokens. k<=0 or k>=vocab disables it."""
if k is None or k <= 0 or k >= logits.size(-1):
return logits
kth = torch.topk(logits, k, dim=-1).values[..., -1, None]
return logits.masked_fill(logits < kth, float("-inf"))
def top_p_filter(logits: torch.Tensor, p: float | None) -> torch.Tensor:
"""
Nucleus sampling: keep the smallest set of tokens whose cumulative
probability reaches p. p>=1.0 disables it (and skips a 32k-element sort).
"""
if p is None or p >= 1.0:
return logits
srt, idx = torch.sort(logits, descending=True, dim=-1)
probs = srt.softmax(dim=-1)
# Subtracting probs shifts the cumulative sum one position right, which
# guarantees the top token is always kept even if it alone exceeds p.
remove = (probs.cumsum(dim=-1) - probs) > p
srt = srt.masked_fill(remove, float("-inf"))
return torch.full_like(logits, float("-inf")).scatter(-1, idx, srt)
# ------------------------------------------------------------------ #
# Combined sampler
# ------------------------------------------------------------------ #
def sample_next_token(
logits: torch.Tensor, # (batch, vocab) raw scores
prev_ids: torch.Tensor, # (batch, n_prev) tokens so far
temperature: float = 0.8, # <= 0.0 selects greedy decoding
top_k: int | None = 50,
top_p: float | None = 1.0,
repetition_penalty: float = 1.0,
generator: torch.Generator | None = None,
) -> torch.Tensor: # (batch, 1) int64
"""Pick the next token. temperature <= 0.0 means deterministic argmax."""
logits = apply_repetition_penalty(
logits.float(), prev_ids, repetition_penalty)
# Greedy. Handled before the division so temperature=0.0 cannot produce
# inf/NaN — the old code divided unconditionally and made greedy decoding
# impossible.
if temperature is None or temperature <= 0.0:
return logits.argmax(dim=-1, keepdim=True)
logits = logits / temperature
logits = top_k_filter(logits, top_k)
logits = top_p_filter(logits, top_p)
probs = torch.softmax(logits, dim=-1)
return torch.multinomial(probs, num_samples=1, generator=generator)
# ------------------------------------------------------------------ #
# Quick self-test
# ------------------------------------------------------------------ #
if __name__ == "__main__":
torch.manual_seed(0)
V = 100
logits = torch.randn(1, V)
prev = torch.tensor([[3, 7, 7]])
# Greedy is deterministic and matches a plain argmax
want = int(logits.argmax())
for temp in (0.0, -1.0, None):
got = sample_next_token(logits, prev, temperature=temp,
repetition_penalty=1.0)
assert int(got) == want, f"greedy failed at temperature={temp}"
print(" greedy decoding: OK")
# top-k restricts the support to exactly k tokens
filtered = top_k_filter(logits.clone(), 5)
assert int(torch.isfinite(filtered).sum()) == 5
assert torch.equal(top_k_filter(logits.clone(), 0), logits), "k=0 disables"
print(" top_k_filter: OK")
# top-p keeps at least one token and never more than the full vocab
for p in (0.01, 0.5, 0.9):
n = int(torch.isfinite(top_p_filter(logits.clone(), p)).sum())
assert 1 <= n <= V, f"top_p={p} kept {n} tokens"
assert torch.equal(top_p_filter(logits.clone(), 1.0), logits), "p=1 disables"
print(" top_p_filter: OK")
# Repetition penalty pushes seen tokens down, leaves others untouched
pen = apply_repetition_penalty(logits.clone(), prev, 2.0)
for t in (3, 7):
assert pen[0, t] < logits[0, t], f"token {t} not penalised"
untouched = [i for i in range(V) if i not in (3, 7)]
assert torch.equal(pen[0, untouched], logits[0, untouched])
print(" repetition_penalty: OK")
# A seeded generator reproduces the same draw
a = sample_next_token(logits, prev, 0.8, 50, 1.0, 1.0,
generator=torch.Generator().manual_seed(42))
b = sample_next_token(logits, prev, 0.8, 50, 1.0, 1.0,
generator=torch.Generator().manual_seed(42))
assert torch.equal(a, b), "seeded sampling is not reproducible"
print(" seeded reproducibility: OK")
print("\nAll sampling tests PASSED.")
|