Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions python/sglang/srt/managers/schedule_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -1192,7 +1192,7 @@ def init_incremental_detokenize(self):

return self.surr_and_decode_ids, self.read_offset - self.surr_offset

def tail_str(self) -> str:
def tail_str(self, new_accepted_len: int = 1) -> str:
# Check stop strings and stop regex patterns together
if (
len(self.sampling_params.stop_strs) == 0
Expand All @@ -1205,7 +1205,12 @@ def tail_str(self) -> str:
self.sampling_params.stop_regex_max_len + 1,
)

tail_len = min(max_len_tail_str, len(self.output_ids))
# Spec decode accepts multiple tokens per step; widen the window to cover
# the whole accepted chunk so a stop string landing mid-chunk (with more
# tokens accepted after it) is not pushed out of view.
tail_len = min(
max_len_tail_str + max(new_accepted_len - 1, 0), len(self.output_ids)
)
return self.tokenizer.decode(self.output_ids[-tail_len:])

def check_match_stop_str_prefix(self) -> bool:
Expand Down Expand Up @@ -1261,12 +1266,12 @@ def _check_token_based_finish(self, new_accepted_tokens: List[int]) -> bool:

return False

def _check_str_based_finish(self):
def _check_str_based_finish(self, new_accepted_len: int = 1):
if (
len(self.sampling_params.stop_strs) > 0
or len(self.sampling_params.stop_regex_strs) > 0
):
tail_str = self.tail_str()
tail_str = self.tail_str(new_accepted_len)

# Check stop strings
if len(self.sampling_params.stop_strs) > 0:
Expand Down Expand Up @@ -1331,7 +1336,7 @@ def update_finish_state(self, new_accepted_len: int = 1):
if self._check_vocab_boundary_finish(new_accepted_tokens):
return

if self._check_str_based_finish():
if self._check_str_based_finish(new_accepted_len):
return

def reset_for_retract(self):
Expand Down
60 changes: 60 additions & 0 deletions test/registered/unit/managers/test_stop_str_speculative.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
"""Regression: under speculative decoding (multi-token commits) a stop string
committed mid-chunk must still trigger the finish check, else the request
over-generates. Drives the real `Req.update_finish_state`; pure CPU."""

import unittest
from array import array

from sglang.srt.managers.schedule_batch import Req
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci

register_cpu_ci(est_time=5, suite="base-a-test-cpu")

STOP_ID = 1
ID_TO_TEXT = {STOP_ID: "STOP", **{i: chr(ord("a") + i % 26) for i in range(10, 40)}}

# "STOP" (index 3) sits 6 tokens back: outside the old (stop_str_max_len + 1)
# window, inside the one widened by new_accepted_len.
MIDCHUNK = [10, 11, 12, STOP_ID, 20, 21, 22, 23, 24]


class _FakeTokenizer:
eos_token_id = -1
additional_stop_token_ids = None

def decode(self, ids):
return "".join(ID_TO_TEXT[int(i)] for i in ids)


def _make_req(output_ids, stop):
sp = SamplingParams(max_new_tokens=1000, stop=stop)
sp.normalize(tokenizer=None) # char-based stop_str_max_len
req = Req(
rid="t",
origin_input_text="",
origin_input_ids=array("q", [0]),
sampling_params=sp,
eos_token_ids=set(),
vocab_size=10_000,
)
req.tokenizer = _FakeTokenizer()
req.output_ids = array("q", output_ids)
return req


class TestStopStrSpeculative(unittest.TestCase):
def test_stop_str_midchunk_finishes(self):
req = _make_req(MIDCHUNK, stop=["STOP"])
req.update_finish_state(new_accepted_len=6)
self.assertTrue(req.finished())
self.assertEqual(req.finished_reason.matched, "STOP")

def test_no_stop_str_does_not_finish(self):
req = _make_req([10, 11, 12, 20, 21, 22, 23, 24], stop=["STOP"])
req.update_finish_state(new_accepted_len=6)
self.assertFalse(req.finished())


if __name__ == "__main__":
unittest.main()
Loading