Fix grammar sync across TP ranks (#17100)

Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
This commit is contained in:
Lianmin Zheng
2026-01-15 18:38:01 -08:00
committed by GitHub
co-authored by Liangsheng Yin
parent c81bad1bf7
commit e7dc85c50b
4 changed files with 112 additions and 116 deletions
@@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import random import time
from concurrent import futures from concurrent import futures
from typing import TYPE_CHECKING, List from typing import TYPE_CHECKING, List
@@ -18,8 +18,6 @@ if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
GRAMMAR_TIMEOUT = envs.SGLANG_GRAMMAR_TIMEOUT.get()
GRAMMAR_POLL_INTERVAL = envs.SGLANG_GRAMMAR_POLL_INTERVAL.get()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -43,6 +41,11 @@ class GrammarManager:
self.grammar_sync_entry = scheduler.dp_tp_group.first_rank self.grammar_sync_entry = scheduler.dp_tp_group.first_rank
self.is_grammar_sync_entry = scheduler.dp_tp_group.is_first_rank self.is_grammar_sync_entry = scheduler.dp_tp_group.is_first_rank
self.SGLANG_GRAMMAR_POLL_INTERVAL = envs.SGLANG_GRAMMAR_POLL_INTERVAL.get()
self.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = (
envs.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS.get()
)
def __len__(self): def __len__(self):
return len(self.grammar_queue) return len(self.grammar_queue)
@@ -101,89 +104,92 @@ class GrammarManager:
return add_to_grammar_queue return add_to_grammar_queue
def finalize_grammar(self, req: Req): def get_ready_grammar_requests(self) -> List[Req]:
self.grammar_backend.set_cache(req.grammar_key, req.grammar.copy()) """
if req.grammar is INVALID_GRAMMAR_OBJ: Move requests whose grammar objects are ready from grammar_queue to waiting_queue.
error_msg = f"Invalid grammar request: {req.grammar_key=}"
Rank i returns two sets ready_reqs_i, failed_reqs_i
ready_reqs_all = all_gather(ready_reqs_i)
failed_reqs_all = all_gather(failed_reqs_i)
ready_reqs = intersect(ready_reqs_all)
failed_reqs = union(failed_reqs_all)
"""
ready_req_idxs: set[int] = set()
failed_req_idxs: set[int] = set()
# Poll for ready requests
start_time = time.perf_counter()
while time.perf_counter() - start_time < self.SGLANG_GRAMMAR_POLL_INTERVAL:
for i, req in enumerate(self.grammar_queue):
if i in ready_req_idxs:
continue
if req.finished() or req.grammar is None: # It is aborted by AbortReq
ready_req_idxs.add(i)
continue
assert isinstance(req.grammar, futures.Future), f"{req=}"
if req.grammar.done():
ready_req_idxs.add(i)
# Sleep a bit to avoid busy waiting
time.sleep(self.SGLANG_GRAMMAR_POLL_INTERVAL / 10)
# Check failed requests
for i, req in enumerate(self.grammar_queue):
if i not in ready_req_idxs:
self.grammar_queue[i].grammar_wait_ct += 1
if (
self.grammar_queue[i].grammar_wait_ct
>= self.SGLANG_GRAMMAR_MAX_POLL_ITERATIONS
):
# Timeout after max poll iterations
# The actual waiting time is SGLANG_GRAMMAR_MAX_POLL_ITERATIONS * max(SGLANG_GRAMMAR_POLL_INTERVAL, GPU_forward_batch_latency)
failed_req_idxs.add(i)
# Sync ready and failed requests across all ranks
if self.grammar_sync_size == 1:
synced_ready_req_idxs = ready_req_idxs
synced_failed_req_idxs = failed_req_idxs
else:
all_gather_output = [None] * self.grammar_sync_size
torch.distributed.all_gather_object(
all_gather_output,
(ready_req_idxs, failed_req_idxs),
group=self.grammar_sync_group,
)
synced_ready_req_idxs = set.intersection(*[x[0] for x in all_gather_output])
synced_failed_req_idxs = set.union(*[x[1] for x in all_gather_output])
# Return ready requests
return_reqs: List[Req] = []
for i in synced_ready_req_idxs:
req = self.grammar_queue[i]
return_reqs.append(req)
if req.finished() or req.grammar is None: # It is aborted by AbortReq
continue
req.grammar = req.grammar.result()
self.grammar_backend.set_cache(req.grammar_key, req.grammar.copy())
if req.grammar is INVALID_GRAMMAR_OBJ:
error_msg = f"Invalid grammar request: {req.grammar_key=}"
req.set_finish_with_abort(error_msg)
# Return failed requests
for i in synced_failed_req_idxs:
req = self.grammar_queue[i]
return_reqs.append(req)
req.grammar.cancel()
self.grammar_backend.set_cache(req.grammar_key, INVALID_GRAMMAR_OBJ)
error_msg = f"Grammar preprocessing timed out: {req.grammar_key=}"
req.set_finish_with_abort(error_msg) req.set_finish_with_abort(error_msg)
def get_ready_grammar_requests(self) -> List[Req]: # Remove finished requests from grammar_queue
"""Move requests whose grammar objects are ready from grammar_queue to waiting_queue.""" self.grammar_queue = [
num_ready_reqs = 0 req
num_timeout_pt = 0 for i, req in enumerate(self.grammar_queue)
sim_timeout_prob = ( if i not in synced_ready_req_idxs and i not in synced_failed_req_idxs
-1 ]
if self.is_grammar_sync_entry return return_reqs
else envs.SGLANG_GRAMMAR_SIMULATE_TIMEOUT.get()
) # Entry rank never simulates timeout
timeout_ct = GRAMMAR_TIMEOUT / GRAMMAR_POLL_INTERVAL
for req in self.grammar_queue:
try:
if req.finished(): # It is aborted by AbortReq
num_ready_reqs += 1
continue
if sim_timeout_prob > 0 and random.random() < sim_timeout_prob:
# Simulate timeout for non-entry ranks in TP sync group for testing
logger.warning(
f"Simulating grammar timeout on {self.scheduler.tp_rank=}"
)
raise futures._base.TimeoutError()
req.grammar = req.grammar.result(timeout=GRAMMAR_POLL_INTERVAL)
self.finalize_grammar(req)
num_ready_reqs += 1
num_timeout_pt = num_ready_reqs
except futures._base.TimeoutError:
req.grammar_wait_ct += 1
# NOTE(lianmin): this timeout is the waiting time of the above line. It is
# not the waiting time from it enters the grammar queue.
if sim_timeout_prob > 0 or req.grammar_wait_ct > timeout_ct:
num_timeout_pt = num_ready_reqs + 1
break
if self.grammar_sync_size > 1:
# Sync across TP ranks to make sure they have the same number of ready requests
tensor = torch.tensor(
[num_ready_reqs, -num_ready_reqs, -num_timeout_pt], dtype=torch.int32
)
torch.distributed.all_reduce(
tensor, op=torch.distributed.ReduceOp.MIN, group=self.grammar_sync_group
)
num_ready_reqs_min = tensor[0].item()
num_ready_reqs_max = -tensor[1].item()
num_timeout_pt_max = -tensor[2].item()
else:
num_ready_reqs_min = num_ready_reqs
num_ready_reqs_max = num_ready_reqs
num_timeout_pt_max = num_timeout_pt
if envs.SGLANG_GRAMMAR_SIMULATE_TIMEOUT.get() > 0:
# NOTE: in simulation timeout mode, if only some TP ranks report a timeout,
# we still treat those requests as ready instead of canceling them.
for i in range(num_ready_reqs_min, num_ready_reqs_max):
req = self.grammar_queue[i]
if req.finished():
continue
if isinstance(req.grammar, futures.Future):
req.grammar = req.grammar.result()
self.finalize_grammar(req)
num_ready_reqs_min = num_ready_reqs_max
# NOTE: in non-simulation mode, cancel any request that times out on a subset of
# TP ranks because timeouts can diverge across ranks.
for i in range(num_ready_reqs_min, num_timeout_pt_max):
req = self.grammar_queue[i]
if req.finished():
continue
if isinstance(req.grammar, futures.Future):
req.grammar.cancel()
req.grammar = INVALID_GRAMMAR_OBJ
self.finalize_grammar(req)
ready_grammar_reqs = self.grammar_queue[:num_timeout_pt_max]
self.grammar_queue = self.grammar_queue[num_timeout_pt_max:]
return ready_grammar_reqs
@@ -212,9 +212,6 @@ class OpenAIServingRerank(OpenAIServingBase):
self._yes_token_id, self._no_token_id = _get_yes_no_token_ids( self._yes_token_id, self._no_token_id = _get_yes_no_token_ids(
tokenizer_manager.tokenizer tokenizer_manager.tokenizer
) )
logger.info(
f"Reranker yes/no token IDs: yes={self._yes_token_id}, no={self._no_token_id}"
)
# NOTE: /v1/rerank is not an official OpenAI endpoint. This module may be moved # NOTE: /v1/rerank is not an official OpenAI endpoint. This module may be moved
# to another module in the future. # to another module in the future.
+2 -6
View File
@@ -172,8 +172,8 @@ class Envs:
SGLANG_TEST_MAX_RETRY = EnvInt(None) SGLANG_TEST_MAX_RETRY = EnvInt(None)
# Constrained Decoding (Grammar) # Constrained Decoding (Grammar)
SGLANG_GRAMMAR_SIMULATE_TIMEOUT = EnvFloat(-1) SGLANG_GRAMMAR_POLL_INTERVAL = EnvFloat(0.005)
SGLANG_GRAMMAR_POLL_INTERVAL = EnvFloat(0.03) SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = EnvInt(10000)
# Test & Debug # Test & Debug
SGLANG_DETECT_SLOW_RANK = EnvBool(False) SGLANG_DETECT_SLOW_RANK = EnvBool(False)
@@ -251,10 +251,6 @@ class Envs:
SGLANG_USE_MESSAGE_QUEUE_BROADCASTER = EnvBool(True) SGLANG_USE_MESSAGE_QUEUE_BROADCASTER = EnvBool(True)
SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS = EnvBool(False) SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS = EnvBool(False)
# Constrained Decoding
SGLANG_DISABLE_OUTLINES_DISK_CACHE = EnvBool(True)
SGLANG_GRAMMAR_TIMEOUT = EnvFloat(300)
# Tool Calling # Tool Calling
SGLANG_FORWARD_UNKNOWN_TOOLS = EnvBool(False) SGLANG_FORWARD_UNKNOWN_TOOLS = EnvBool(False)
@@ -3,7 +3,6 @@ from types import SimpleNamespace
import requests import requests
from sglang.srt.environ import envs
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
@@ -35,18 +34,17 @@ class TestDPAttentionDP2TP4(
def setUpClass(cls): def setUpClass(cls):
cls.model = DEFAULT_MLA_MODEL_NAME_FOR_TEST cls.model = DEFAULT_MLA_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST cls.base_url = DEFAULT_URL_FOR_TEST
with envs.SGLANG_GRAMMAR_SIMULATE_TIMEOUT.override(0.5): cls.process = popen_launch_server(
cls.process = popen_launch_server( cls.model,
cls.model, cls.base_url,
cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=[
other_args=[ "--trust-remote-code",
"--trust-remote-code", "--tp=4",
"--tp=4", "--enable-dp-attention",
"--enable-dp-attention", "--dp=2",
"--dp=2", ],
], )
)
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
@@ -91,13 +89,12 @@ class TestDPAttentionDP2TP2DeepseekV3MTP(
] ]
if not is_in_amd_ci(): if not is_in_amd_ci():
other_args += ["--mem-frac", "0.7"] other_args += ["--mem-frac", "0.7"]
with envs.SGLANG_GRAMMAR_SIMULATE_TIMEOUT.override(0.5): cls.process = popen_launch_server(
cls.process = popen_launch_server( cls.model,
cls.model, cls.base_url,
cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=other_args,
other_args=other_args, )
)
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):