Decouple grammar logic out of scheduler. (#16820)

This commit is contained in:
Liangsheng Yin
2026-01-12 10:52:42 +08:00
committed by GitHub
parent 38a88479c6
commit 5b7bed7ca4
6 changed files with 186 additions and 143 deletions
@@ -0,0 +1,163 @@
from __future__ import annotations
import logging
from concurrent import futures
from typing import TYPE_CHECKING, List
import torch
from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
create_grammar_backend,
)
from sglang.srt.environ import envs
if TYPE_CHECKING:
from sglang.srt.managers.io_struct import AbortReq
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import Scheduler
GRAMMAR_TIMEOUT = envs.SGLANG_GRAMMAR_TIMEOUT.get()
logger = logging.getLogger(__name__)
class GrammarManager:
def __init__(self, scheduler: Scheduler):
self.scheduler = scheduler
self.server_args = scheduler.server_args
self.grammar_queue: List[Req] = []
if not self.server_args.skip_tokenizer_init:
self.grammar_backend = create_grammar_backend(
self.server_args,
scheduler.tokenizer,
scheduler.model_config.vocab_size,
scheduler.model_config.hf_eos_token_id,
)
else:
self.grammar_backend = None
def __len__(self):
return len(self.grammar_queue)
def clear(self):
if self.grammar_backend:
self.grammar_backend.reset()
def has_waiting_grammars(self) -> bool:
return len(self.grammar_queue) > 0
def abort_requests(self, recv_req: AbortReq):
for req in self.grammar_queue:
if recv_req.abort_all or req.rid.startswith(recv_req.rid):
logger.debug(f"Abort grammar queue request. {req.rid=}")
if req.grammar:
req.grammar.cancel()
req.set_finish_with_abort("Aborted by AbortReq.")
def process_req_with_grammar(self, req: Req) -> bool:
# Init grammar cache for this request
add_to_grammar_queue = False
if (
req.sampling_params.json_schema is not None
or req.sampling_params.regex is not None
or req.sampling_params.ebnf is not None
or req.sampling_params.structural_tag is not None
):
if self.grammar_backend is None:
error_msg = "Grammar-based generation (json_schema, regex, ebnf, structural_tag) is not supported when the server is launched with --grammar-backend none"
req.set_finish_with_abort(error_msg)
else:
if req.sampling_params.json_schema is not None:
key = ("json", req.sampling_params.json_schema)
elif req.sampling_params.regex is not None:
key = ("regex", req.sampling_params.regex)
elif req.sampling_params.ebnf is not None:
key = ("ebnf", req.sampling_params.ebnf)
elif req.sampling_params.structural_tag:
key = ("structural_tag", req.sampling_params.structural_tag)
value, cache_hit = self.grammar_backend.get_cached_or_future_value(
key, req.require_reasoning
)
req.grammar = value
if not cache_hit:
req.grammar_key = key
add_to_grammar_queue = True
else:
if value is INVALID_GRAMMAR_OBJ: # We hit a cached invalid grammar.
error_msg = f"Invalid grammar request with cache hit: {key=}"
req.set_finish_with_abort(error_msg)
if add_to_grammar_queue:
self.grammar_queue.append(req)
return add_to_grammar_queue
def get_ready_grammar_requests(self) -> List[Req]:
"""Move requests whose grammar objects are ready from grammar_queue to waiting_queue."""
num_ready_reqs = 0
num_timeout_reqs = 0
for req in self.grammar_queue:
try:
if req.finished(): # It is aborted by AbortReq
num_ready_reqs += 1
continue
req.grammar = req.grammar.result(timeout=0.03)
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)
num_ready_reqs += 1
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 req.grammar_wait_ct > GRAMMAR_TIMEOUT / 0.03:
num_timeout_reqs = 1
break
if self.server_args.enable_dp_attention:
tp_size = self.scheduler.attn_tp_size
tp_group = self.scheduler.attn_tp_cpu_group
else:
tp_size = self.scheduler.tp_size
tp_group = self.scheduler.tp_cpu_group
if tp_size > 1:
# Sync across TP ranks to make sure they have the same number of ready requests
tensor = torch.tensor([num_ready_reqs, num_timeout_reqs], dtype=torch.int32)
torch.distributed.all_reduce(
tensor, op=torch.distributed.ReduceOp.MAX, group=tp_group
)
num_ready_reqs_max, num_timeout_reqs_max = tensor.tolist()
for i in range(num_ready_reqs, num_ready_reqs_max):
req = self.grammar_queue[i]
if req.finished(): # 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)
else:
num_ready_reqs_max = num_ready_reqs
num_timeout_reqs_max = num_timeout_reqs
for i in range(num_ready_reqs, num_ready_reqs + num_timeout_reqs_max):
req = self.grammar_queue[i]
req.grammar.cancel()
self.grammar_backend.set_cache(req.grammar_key, INVALID_GRAMMAR_OBJ)
error_msg = f"Grammar preprocessing timed out for {req.grammar_key=}"
req.set_finish_with_abort(error_msg)
num_ready_reqs = num_ready_reqs_max + num_timeout_reqs_max
ready_grammar_reqs = self.grammar_queue[:num_ready_reqs]
self.grammar_queue = self.grammar_queue[num_ready_reqs:]
return ready_grammar_reqs
+4 -2
View File
@@ -946,8 +946,10 @@ class SchedulerDisaggregationDecodeMixin:
def get_new_prebuilt_batch(self: Scheduler) -> Optional[ScheduleBatch]: def get_new_prebuilt_batch(self: Scheduler) -> Optional[ScheduleBatch]:
"""Create a schedulebatch for fake completed prefill""" """Create a schedulebatch for fake completed prefill"""
if self.grammar_queue: if self.grammar_manager.has_waiting_grammars():
self.move_ready_grammar_requests() ready_grammar_requests = self.grammar_manager.get_ready_grammar_requests()
for req in ready_grammar_requests:
self._add_request_to_queue(req)
if len(self.waiting_queue) == 0: if len(self.waiting_queue) == 0:
return None return None
@@ -719,6 +719,7 @@ class Req:
self.embedding = None self.embedding = None
# Constrained decoding # Constrained decoding
self.grammar_key: Optional[str] = None
self.grammar: Optional[BaseGrammarObject] = None self.grammar: Optional[BaseGrammarObject] = None
self.grammar_wait_ct = 0 self.grammar_wait_ct = 0
+12 -135
View File
@@ -20,7 +20,6 @@ import signal
import sys import sys
import time import time
from collections import deque from collections import deque
from concurrent import futures
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import Any, Deque, Dict, List, Optional, Tuple, Union from typing import Any, Deque, Dict, List, Optional, Tuple, Union
@@ -35,10 +34,7 @@ from torch.cuda import StreamContext as CudaStreamContext
from torch.distributed import barrier from torch.distributed import barrier
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.constrained.base_grammar_backend import ( from sglang.srt.constrained.grammar_manager import GrammarManager
INVALID_GRAMMAR_OBJ,
create_grammar_backend,
)
from sglang.srt.disaggregation.decode import ( from sglang.srt.disaggregation.decode import (
DecodePreallocQueue, DecodePreallocQueue,
DecodeTransferQueue, DecodeTransferQueue,
@@ -212,7 +208,6 @@ logger = logging.getLogger(__name__)
TEST_RETRACT = envs.SGLANG_TEST_RETRACT.get() TEST_RETRACT = envs.SGLANG_TEST_RETRACT.get()
TEST_RETRACT_INTERVAL = envs.SGLANG_TEST_RETRACT_INTERVAL.get() TEST_RETRACT_INTERVAL = envs.SGLANG_TEST_RETRACT_INTERVAL.get()
TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get() TEST_RETRACT_NO_PREFILL_BS = envs.SGLANG_TEST_RETRACT_NO_PREFILL_BS.get()
GRAMMAR_TIMEOUT = float(os.environ.get("SGLANG_GRAMMAR_TIMEOUT", 300))
@dataclass @dataclass
@@ -355,9 +350,6 @@ class Scheduler(
# Init chunked prefill # Init chunked prefill
self.init_chunked_prefill() self.init_chunked_prefill()
# Init the grammar backend for constrained generation
self.init_grammar_backend()
# Init schedule policy and new token estimation # Init schedule policy and new token estimation
self.init_schedule_policy() self.init_schedule_policy()
@@ -379,6 +371,9 @@ class Scheduler(
# Init request dispatcher # Init request dispatcher
self.init_request_dispatcher() self.init_request_dispatcher()
# Init the grammar backend for constrained generation
self.grammar_manager = GrammarManager(self)
self.is_initializing = False self.is_initializing = False
def init_model_config(self): def init_model_config(self):
@@ -756,18 +751,6 @@ class Scheduler(
) )
self.enable_dynamic_chunking = False self.enable_dynamic_chunking = False
def init_grammar_backend(self):
self.grammar_queue: List[Req] = []
if not self.server_args.skip_tokenizer_init:
self.grammar_backend = create_grammar_backend(
self.server_args,
self.tokenizer,
self.model_config.vocab_size,
self.model_config.hf_eos_token_id,
)
else:
self.grammar_backend = None
def init_schedule_policy(self): def init_schedule_policy(self):
# Init schedule policy and new token estimation # Init schedule policy and new token estimation
self.policy = SchedulePolicy( self.policy = SchedulePolicy(
@@ -1554,43 +1537,8 @@ class Scheduler(
self._add_request_to_queue(req) self._add_request_to_queue(req)
return return
# Init grammar cache for this request added_to_grammar_queue = self.grammar_manager.process_req_with_grammar(req)
add_to_grammar_queue = False if not added_to_grammar_queue:
if (
req.sampling_params.json_schema is not None
or req.sampling_params.regex is not None
or req.sampling_params.ebnf is not None
or req.sampling_params.structural_tag is not None
):
if self.grammar_backend is None:
error_msg = "Grammar-based generation (json_schema, regex, ebnf, structural_tag) is not supported when the server is launched with --grammar-backend none"
req.set_finish_with_abort(error_msg)
else:
if req.sampling_params.json_schema is not None:
key = ("json", req.sampling_params.json_schema)
elif req.sampling_params.regex is not None:
key = ("regex", req.sampling_params.regex)
elif req.sampling_params.ebnf is not None:
key = ("ebnf", req.sampling_params.ebnf)
elif req.sampling_params.structural_tag:
key = ("structural_tag", req.sampling_params.structural_tag)
value, cache_hit = self.grammar_backend.get_cached_or_future_value(
key, req.require_reasoning
)
req.grammar = value
if not cache_hit:
req.grammar_key = key
add_to_grammar_queue = True
else:
if value is INVALID_GRAMMAR_OBJ: # We hit a cached invalid grammar.
error_msg = f"Invalid grammar request with cache hit: {key=}"
req.set_finish_with_abort(error_msg)
if add_to_grammar_queue:
self.grammar_queue.append(req)
else:
self._add_request_to_queue(req) self._add_request_to_queue(req)
def handle_batch_generate_request( def handle_batch_generate_request(
@@ -1889,8 +1837,10 @@ class Scheduler(
self, prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor] self, prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor]
) -> Optional[ScheduleBatch]: ) -> Optional[ScheduleBatch]:
# Check if the grammar is ready in the grammar queue # Check if the grammar is ready in the grammar queue
if self.grammar_queue: if self.grammar_manager.has_waiting_grammars():
self.move_ready_grammar_requests() ready_grammar_requests = self.grammar_manager.get_ready_grammar_requests()
for req in ready_grammar_requests:
self._add_request_to_queue(req)
if self.try_preemption: if self.try_preemption:
# Reset batch_is_full to try preemption with a prefill adder. # Reset batch_is_full to try preemption with a prefill adder.
@@ -2369,73 +2319,6 @@ class Scheduler(
self.return_health_check_ct -= 1 self.return_health_check_ct -= 1
self.send_to_tokenizer.send_output(HealthCheckOutput()) self.send_to_tokenizer.send_output(HealthCheckOutput())
def move_ready_grammar_requests(self):
"""Move requests whose grammar objects are ready from grammar_queue to waiting_queue."""
num_ready_reqs = 0
num_timeout_reqs = 0
for req in self.grammar_queue:
try:
if req.finished(): # It is aborted by AbortReq
num_ready_reqs += 1
continue
req.grammar = req.grammar.result(timeout=0.03)
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)
num_ready_reqs += 1
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 req.grammar_wait_ct > GRAMMAR_TIMEOUT / 0.03:
num_timeout_reqs = 1
break
if self.server_args.enable_dp_attention:
tp_size = self.attn_tp_size
tp_group = self.attn_tp_cpu_group
else:
tp_size = self.tp_size
tp_group = self.tp_cpu_group
if tp_size > 1:
# Sync across TP ranks to make sure they have the same number of ready requests
tensor = torch.tensor([num_ready_reqs, num_timeout_reqs], dtype=torch.int32)
torch.distributed.all_reduce(
tensor, op=torch.distributed.ReduceOp.MAX, group=tp_group
)
num_ready_reqs_max, num_timeout_reqs_max = tensor.tolist()
for i in range(num_ready_reqs, num_ready_reqs_max):
req = self.grammar_queue[i]
if req.finished(): # 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)
else:
num_ready_reqs_max = num_ready_reqs
num_timeout_reqs_max = num_timeout_reqs
for i in range(num_ready_reqs, num_ready_reqs + num_timeout_reqs_max):
req = self.grammar_queue[i]
req.grammar.cancel()
self.grammar_backend.set_cache(req.grammar_key, INVALID_GRAMMAR_OBJ)
error_msg = f"Grammar preprocessing timed out for {req.grammar_key=}"
req.set_finish_with_abort(error_msg)
num_ready_reqs = num_ready_reqs_max + num_timeout_reqs_max
for req in self.grammar_queue[:num_ready_reqs]:
self._add_request_to_queue(req)
self.grammar_queue = self.grammar_queue[num_ready_reqs:]
def flush_cache_wrapped(self, recv_req: FlushCacheReqInput): def flush_cache_wrapped(self, recv_req: FlushCacheReqInput):
success = self.flush_cache() success = self.flush_cache()
return FlushCacheReqOutput(success=success) return FlushCacheReqOutput(success=success)
@@ -2478,8 +2361,7 @@ class Scheduler(
self.tree_cache.reset() self.tree_cache.reset()
self.req_to_token_pool.clear() self.req_to_token_pool.clear()
self.token_to_kv_pool_allocator.clear() self.token_to_kv_pool_allocator.clear()
if self.grammar_backend: self.grammar_manager.clear()
self.grammar_backend.reset()
self.reset_metrics() self.reset_metrics()
if self.draft_worker: if self.draft_worker:
@@ -2616,15 +2498,10 @@ class Scheduler(
logger.debug(f"Abort queued request. {req.rid=}") logger.debug(f"Abort queued request. {req.rid=}")
# Delete the requests in the grammar queue # Delete the requests in the grammar queue
for req in self.grammar_queue:
# Abort method 2: call `set_finish_with_abort` # Abort method 2: call `set_finish_with_abort`
# The request will still run one prefill forward pass. # The request will still run one prefill forward pass.
# In this case, we change the input_ids to be only one token to make this prefill cheap. # In this case, we change the input_ids to be only one token to make this prefill cheap.
if recv_req.abort_all or req.rid.startswith(recv_req.rid): self.grammar_manager.abort_requests(recv_req)
logger.debug(f"Abort grammar queue request. {req.rid=}")
if req.grammar:
req.grammar.cancel()
req.set_finish_with_abort("Aborted by AbortReq.")
# Delete requests not in the waiting queue when PD disaggregation is enabled # Delete requests not in the waiting queue when PD disaggregation is enabled
if self.disaggregation_mode == DisaggregationMode.PREFILL: if self.disaggregation_mode == DisaggregationMode.PREFILL:
@@ -227,7 +227,7 @@ class SchedulerMetricsMixin:
if self.is_hybrid_ssm: if self.is_hybrid_ssm:
self.stats.mamba_usage = mamba_usage self.stats.mamba_usage = mamba_usage
self.stats.num_queue_reqs = len(self.waiting_queue) self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_grammar_queue_reqs = len(self.grammar_queue) self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
self.stats.cache_hit_rate = cache_hit_rate self.stats.cache_hit_rate = cache_hit_rate
self.stats.max_total_num_tokens = self.max_total_num_tokens self.stats.max_total_num_tokens = self.max_total_num_tokens
@@ -393,7 +393,7 @@ class SchedulerMetricsMixin:
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item() self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
self.stats.gen_throughput = self.last_gen_throughput self.stats.gen_throughput = self.last_gen_throughput
self.stats.num_queue_reqs = len(self.waiting_queue) self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_grammar_queue_reqs = len(self.grammar_queue) self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
self.stats.cache_hit_rate = cache_hit_rate self.stats.cache_hit_rate = cache_hit_rate
self.stats.max_total_num_tokens = self.max_total_num_tokens self.stats.max_total_num_tokens = self.max_total_num_tokens
@@ -288,7 +288,7 @@ class SchedulerRuntimeCheckerMixin:
self.stats.token_usage = round(token_usage, 2) self.stats.token_usage = round(token_usage, 2)
self.stats.gen_throughput = 0 self.stats.gen_throughput = 0
self.stats.num_queue_reqs = len(self.waiting_queue) self.stats.num_queue_reqs = len(self.waiting_queue)
self.stats.num_grammar_queue_reqs = len(self.grammar_queue) self.stats.num_grammar_queue_reqs = len(self.grammar_manager)
if self.disaggregation_mode == DisaggregationMode.PREFILL: if self.disaggregation_mode == DisaggregationMode.PREFILL:
self.stats.num_prefill_prealloc_queue_reqs = len( self.stats.num_prefill_prealloc_queue_reqs = len(
self.disagg_prefill_bootstrap_queue.queue self.disagg_prefill_bootstrap_queue.queue