From 556cf54d4783674b14cdff567aac6fd9b6bedcd8 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Mon, 15 Jun 2026 23:32:28 -0700 Subject: [PATCH] [Perf] Avoid per-decode-step host sync in min_new_tokens penalty (#28397) --- python/sglang/srt/sampling/penaltylib/min_new_tokens.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) 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]