diff --git a/python/sglang/srt/sampling/penaltylib/min_new_tokens.py b/python/sglang/srt/sampling/penaltylib/min_new_tokens.py index 08f76e1f1..692f226e4 100644 --- a/python/sglang/srt/sampling/penaltylib/min_new_tokens.py +++ b/python/sglang/srt/sampling/penaltylib/min_new_tokens.py @@ -67,8 +67,11 @@ class BatchedMinNewTokensPenalizer(_BatchedPenalizer): self.len_output_tokens += 1 def _apply(self, logits: torch.Tensor): - mask = (self.len_output_tokens < self.min_new_tokens).expand_as(logits) - logits[mask] += self.stop_token_penalties[mask] + # Boolean-mask indexing (logits[mask]) is data-dependent and forces a + # device-to-host sync every decode step; torch.where is a plain + # elementwise select with no sync (and no -inf*0=nan). + mask = self.len_output_tokens < self.min_new_tokens + logits.add_(torch.where(mask, self.stop_token_penalties, 0.0)) def _filter(self, keep_indices: torch.Tensor): self.min_new_tokens = self.min_new_tokens[keep_indices]