Download code/alpamayo_tt/sampling.py from changh95/Alpamayo2-Super-p300x2: direct link, hf CLI and curl.
- Browser
- Download file 2.1 kB
-
https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/sampling.py
- Command line
-
hf download hf://changh95/Alpamayo2-Super-p300x2/code/alpamayo_tt/sampling.py
-
curl -L -o sampling.py https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/sampling.py
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 | |