diff --git a/python/freetoken/core.py b/python/freetoken/core.py index 269b94ffe9..739f7f994d 100644 --- a/python/freetoken/core.py +++ b/python/freetoken/core.py @@ -27,7 +27,7 @@ class SamplingParams: @property def is_greedy(self) -> bool: - return (self.temperature <= 0.0 or self.top_k == 1) and self.top_p == 1.0 + return self.temperature <= 0.0 or self.top_k == 1 @dataclass(eq=False) diff --git a/python/freetoken/engine/sample.py b/python/freetoken/engine/sample.py index 01d14b1aab..3bdae76957 100644 --- a/python/freetoken/engine/sample.py +++ b/python/freetoken/engine/sample.py @@ -15,6 +15,7 @@ class BatchSamplingArgs: temperatures: torch.Tensor | None top_k: torch.Tensor | None = None top_p: torch.Tensor | None = None + greedy_mask: torch.Tensor | None = None def make_device_tensor(data: List, dtype: torch.dtype, device: torch.device) -> torch.Tensor: @@ -57,24 +58,42 @@ class Sampler: def prepare(self, batch: Batch) -> BatchSamplingArgs: params = [r.sampling_params for r in batch.reqs] - if all(p.is_greedy for p in params): + is_greedy = [p.is_greedy for p in params] + if all(is_greedy): return BatchSamplingArgs(temperatures=None) MIN_P = MIN_T = 1e-6 - ts = [max(0.0 if p.is_greedy else p.temperature, MIN_T) for p in params] - top_ks = [p.top_k if p.top_k >= 1 else self.vocab_size for p in params] - top_ps = [min(max(p.top_p, MIN_P), 1.0) for p in params] + # Greedy outputs are selected explicitly in sample(); use neutral sampling + # parameters for those rows instead of approximating argmax at low temperature. + ts = [1.0 if g else max(p.temperature, MIN_T) for p, g in zip(params, is_greedy)] + top_ks = [ + p.top_k if not g and p.top_k >= 1 else self.vocab_size + for p, g in zip(params, is_greedy) + ] + top_ps = [ + 1.0 if g else min(max(p.top_p, MIN_P), 1.0) + for p, g in zip(params, is_greedy) + ] temperatures = make_device_tensor(ts, torch.float32, self.device) top_k, top_p = None, None if any(k != self.vocab_size for k in top_ks): top_k = make_device_tensor(top_ks, torch.int32, self.device) if any(p < 1.0 for p in top_ps): top_p = make_device_tensor(top_ps, torch.float32, self.device) - return BatchSamplingArgs(temperatures, top_k=top_k, top_p=top_p) + greedy_mask = ( + make_device_tensor(is_greedy, torch.bool, self.device) if any(is_greedy) else None + ) + return BatchSamplingArgs(temperatures, top_k=top_k, top_p=top_p, greedy_mask=greedy_mask) @nvtx_annotate("Sampler") def sample(self, logits: torch.Tensor, args: BatchSamplingArgs) -> torch.Tensor: with torch.cuda.nvtx.range("Sampler"): if args.temperatures is None: # greedy sampling return torch.argmax(logits, dim=-1) - return sample_impl(logits.float(), args.temperatures, args.top_k, args.top_p) + tokens = sample_impl(logits.float(), args.temperatures, args.top_k, args.top_p) + if args.greedy_mask is not None: + # Mixed batches still run probability sampling for all rows, but + # greedy rows must follow argmax's deterministic tie-breaking. + greedy_tokens = torch.argmax(logits, dim=-1).to(tokens.dtype) + tokens = torch.where(args.greedy_mask, greedy_tokens, tokens) + return tokens