diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 42ea84310..c85fcacad 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1192,7 +1192,7 @@ class Req(ReqDllmMixin): 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 @@ -1205,7 +1205,12 @@ class Req(ReqDllmMixin): 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: @@ -1261,12 +1266,12 @@ class Req(ReqDllmMixin): 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: @@ -1331,7 +1336,7 @@ class Req(ReqDllmMixin): 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): diff --git a/test/registered/unit/managers/test_stop_str_speculative.py b/test/registered/unit/managers/test_stop_str_speculative.py new file mode 100644 index 000000000..ce1bf4fcb --- /dev/null +++ b/test/registered/unit/managers/test_stop_str_speculative.py @@ -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()