Fix grammar sync across TP ranks (#17100)
Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
This commit is contained in:
co-authored by
Liangsheng Yin
parent
c81bad1bf7
commit
e7dc85c50b
@@ -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.
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user