"""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