fix: stop-string check misses early matches during speculative decoding (#23802)
Co-authored-by: xythink <xythink@users.noreply.github.com> Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
co-authored by
xythink
hnyls2002
Liangsheng Yin
parent
991689fd0d
commit
d145a6127a
@@ -1192,7 +1192,7 @@ class Req(ReqDllmMixin):
|
|||||||
|
|
||||||
return self.surr_and_decode_ids, self.read_offset - self.surr_offset
|
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
|
# Check stop strings and stop regex patterns together
|
||||||
if (
|
if (
|
||||||
len(self.sampling_params.stop_strs) == 0
|
len(self.sampling_params.stop_strs) == 0
|
||||||
@@ -1205,7 +1205,12 @@ class Req(ReqDllmMixin):
|
|||||||
self.sampling_params.stop_regex_max_len + 1,
|
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:])
|
return self.tokenizer.decode(self.output_ids[-tail_len:])
|
||||||
|
|
||||||
def check_match_stop_str_prefix(self) -> bool:
|
def check_match_stop_str_prefix(self) -> bool:
|
||||||
@@ -1261,12 +1266,12 @@ class Req(ReqDllmMixin):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _check_str_based_finish(self):
|
def _check_str_based_finish(self, new_accepted_len: int = 1):
|
||||||
if (
|
if (
|
||||||
len(self.sampling_params.stop_strs) > 0
|
len(self.sampling_params.stop_strs) > 0
|
||||||
or len(self.sampling_params.stop_regex_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
|
# Check stop strings
|
||||||
if len(self.sampling_params.stop_strs) > 0:
|
if len(self.sampling_params.stop_strs) > 0:
|
||||||
@@ -1331,7 +1336,7 @@ class Req(ReqDllmMixin):
|
|||||||
if self._check_vocab_boundary_finish(new_accepted_tokens):
|
if self._check_vocab_boundary_finish(new_accepted_tokens):
|
||||||
return
|
return
|
||||||
|
|
||||||
if self._check_str_based_finish():
|
if self._check_str_based_finish(new_accepted_len):
|
||||||
return
|
return
|
||||||
|
|
||||||
def reset_for_retract(self):
|
def reset_for_retract(self):
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user