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:
Jing Wang
2026-07-12 02:54:50 +08:00
committed by GitHub
co-authored by ziang663 Chao Shi
parent 9b4bb415dd
commit ed554aac17
6 changed files with 239 additions and 48 deletions
@@ -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."""