changh95's picture
Add files using upload-large-folder tool
0190e6b verified
Raw History Blame Contribute Delete
2.1 kB
"""Host-side next-token selection for the chain-of-causation decode.
Mirrors the reference: discrete trajectory-token logits and the text EOS ids are masked to -inf,
then either greedy argmax (parity tests) or temperature + top-p sampling (default top_p 0.98, T 0.6).
Logits arrive as float32 [V] (or [1, V]) from the device.
"""
from __future__ import annotations
import torch
from .config import ModelConfig
class TokenSampler:
def __init__(self, cfg: ModelConfig, top_p: float = 0.98, temperature: float = 0.6, greedy: bool = False,
seed: int | None = None, top_k: int | None = None):
self.cfg = cfg
self.top_p = top_p
self.temperature = temperature
self.greedy = greedy
self.top_k = top_k
self.gen = torch.Generator()
if seed is not None:
self.gen.manual_seed(seed)
self.mask = torch.zeros(cfg.text.vocab, dtype=torch.bool)
for lo, hi in cfg.masked_logit_ranges:
self.mask[lo:hi] = True
self.stop_id = cfg.traj.future_start
def __call__(self, logits: torch.Tensor) -> int:
logits = logits.reshape(-1).float().clone()
logits[self.mask] = float("-inf")
if self.greedy or self.temperature <= 0:
return int(torch.argmax(logits))
logits = logits / self.temperature
if self.top_k:
kth = torch.topk(logits, self.top_k).values[-1]
logits[logits < kth] = float("-inf")
probs = torch.softmax(logits, dim=-1)
if self.top_p < 1.0:
sp, si = torch.sort(probs, descending=True)
cum = torch.cumsum(sp, dim=-1)
# HF TopPLogitsWarper keeps the smallest set whose cumulative prob >= top_p (min 1 token)
remove = cum - sp > self.top_p
sp[remove] = 0.0
sp = sp / sp.sum()
idx = torch.multinomial(sp, 1, generator=self.gen)
return int(si[idx])
return int(torch.multinomial(probs, 1, generator=self.gen))
def is_stop(self, token: int) -> bool:
return token == self.stop_id