Fix: add grammar sync in PP for structured output (#30747)
Signed-off-by: Jing Wang <jingwang96@qq.com> Co-authored-by: ziang663 <119752791+ziang663@users.noreply.github.com> Co-authored-by: Chao Shi <chao.shi@alibaba-inc.com>
This commit is contained in:
co-authored by
ziang663
Chao Shi
parent
9b4bb415dd
commit
ed554aac17
@@ -25,6 +25,7 @@ from sglang.srt.constrained.base_grammar_backend import (
|
||||
)
|
||||
from sglang.srt.constrained.grammar_manager import GrammarManager
|
||||
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
|
||||
from sglang.srt.distributed.communication_tags import P2PTag
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(2.0, "base-a-test-cpu")
|
||||
@@ -45,6 +46,9 @@ def _make_scheduler(grammar_backend_name="none", skip_tokenizer=False):
|
||||
scheduler.dp_tp_group.world_size = 1
|
||||
scheduler.dp_tp_group.first_rank = 0
|
||||
scheduler.dp_tp_group.is_first_rank = True
|
||||
scheduler.ps.pp_rank = 0
|
||||
scheduler.ps.pp_size = 1
|
||||
scheduler.pp_group = None
|
||||
|
||||
return scheduler
|
||||
|
||||
@@ -627,6 +631,92 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
|
||||
self.assertEqual(len(mgr.grammar_queue), 0)
|
||||
|
||||
|
||||
class _FakePPSendWork:
|
||||
def __init__(self):
|
||||
self.waited = False
|
||||
self.work = self
|
||||
|
||||
def wait(self):
|
||||
self.waited = True
|
||||
|
||||
|
||||
class _FakePPGroup:
|
||||
def __init__(self, recv_data=None):
|
||||
self.recv_data = recv_data
|
||||
self.recv_calls = []
|
||||
self.send_calls = []
|
||||
|
||||
def recv_object(self, *, src, tag):
|
||||
self.recv_calls.append((src, tag))
|
||||
return self.recv_data
|
||||
|
||||
def send_object(self, data, *, dst, async_send, tag):
|
||||
self.send_calls.append((data, dst, async_send, tag))
|
||||
return [_FakePPSendWork()]
|
||||
|
||||
|
||||
class TestGrammarManagerPPSync(unittest.TestCase):
|
||||
"""Test PP synchronization of grammar ready/failed indexes."""
|
||||
|
||||
def _make_mgr_for_pp(self, pp_rank, pp_size, pp_group):
|
||||
scheduler = _make_scheduler()
|
||||
scheduler.server_args.skip_tokenizer_init = True
|
||||
scheduler.ps.pp_rank = pp_rank
|
||||
scheduler.ps.pp_size = pp_size
|
||||
scheduler.pp_group = pp_group
|
||||
mgr = GrammarManager(scheduler)
|
||||
mgr.grammar_backend = MagicMock(spec=BaseGrammarBackend)
|
||||
return mgr
|
||||
|
||||
def test_pp0_sends_ready_failed_without_recv(self):
|
||||
pp_group = _FakePPGroup()
|
||||
mgr = self._make_mgr_for_pp(pp_rank=0, pp_size=3, pp_group=pp_group)
|
||||
|
||||
data = mgr._pp_sync_ready_failed({1}, {3})
|
||||
|
||||
self.assertEqual(data, ({1}, {3}))
|
||||
self.assertEqual(pp_group.recv_calls, [])
|
||||
self.assertEqual(
|
||||
pp_group.send_calls,
|
||||
[(({1}, {3}), 1, True, P2PTag.GRAMMAR_PP_SYNC)],
|
||||
)
|
||||
|
||||
def test_middle_pp_rank_receives_and_forwards_pp0_result(self):
|
||||
pp0_data = ({1, 2}, {4})
|
||||
pp_group = _FakePPGroup(recv_data=pp0_data)
|
||||
mgr = self._make_mgr_for_pp(pp_rank=1, pp_size=3, pp_group=pp_group)
|
||||
|
||||
data = mgr._pp_sync_ready_failed(set(), set())
|
||||
|
||||
self.assertEqual(data, pp0_data)
|
||||
self.assertEqual(pp_group.recv_calls, [(0, P2PTag.GRAMMAR_PP_SYNC)])
|
||||
self.assertEqual(
|
||||
pp_group.send_calls,
|
||||
[(pp0_data, 2, True, P2PTag.GRAMMAR_PP_SYNC)],
|
||||
)
|
||||
|
||||
def test_last_pp_rank_receives_without_forwarding(self):
|
||||
pp0_data = ({0}, {2})
|
||||
pp_group = _FakePPGroup(recv_data=pp0_data)
|
||||
mgr = self._make_mgr_for_pp(pp_rank=2, pp_size=3, pp_group=pp_group)
|
||||
|
||||
data = mgr._pp_sync_ready_failed(set(), set())
|
||||
|
||||
self.assertEqual(data, pp0_data)
|
||||
self.assertEqual(pp_group.recv_calls, [(1, P2PTag.GRAMMAR_PP_SYNC)])
|
||||
self.assertEqual(pp_group.send_calls, [])
|
||||
|
||||
def test_pp_sync_drains_previous_async_send_work(self):
|
||||
pp_group = _FakePPGroup()
|
||||
mgr = self._make_mgr_for_pp(pp_rank=0, pp_size=2, pp_group=pp_group)
|
||||
work = _FakePPSendWork()
|
||||
mgr.grammar_pp_sync_work_list = [work]
|
||||
|
||||
mgr._pp_sync_ready_failed({1}, set())
|
||||
|
||||
self.assertTrue(work.waited)
|
||||
|
||||
|
||||
class TestStrictReasoningPaths(unittest.TestCase):
|
||||
"""Test _enable_strict_thinking code paths in GrammarManager."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user