Minor style fixes to the scheduler.py (#15218)

This commit is contained in:
Lianmin Zheng
2025-12-16 17:09:44 -08:00
committed by GitHub
parent 46ad4b986d
commit 9d64a7b24f
11 changed files with 207 additions and 204 deletions
+1 -4
View File
@@ -228,10 +228,7 @@ class Engine(EngineBase):
) )
self.tokenizer_manager = tokenizer_manager self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager self.template_manager = template_manager
self.scheduler_info = scheduler_infos[0]
scheduler_info = scheduler_infos[0]
self.scheduler_info = scheduler_info
self.port_args = port_args self.port_args = port_args
self.remote_instance_transfer_engine_info = ( self.remote_instance_transfer_engine_info = (
parse_remote_instance_transfer_engine_info_from_scheduler_infos( parse_remote_instance_transfer_engine_info_from_scheduler_infos(
+1 -1
View File
@@ -1432,7 +1432,7 @@ def _execute_server_warmup(
for _ in range(120): for _ in range(120):
time.sleep(1) time.sleep(1)
try: try:
res = requests.get(url + "/get_model_info", timeout=5, headers=headers) res = requests.get(url + "/model_info", timeout=5, headers=headers)
assert res.status_code == 200, f"{res=}, {res.text=}" assert res.status_code == 200, f"{res=}, {res.text=}"
success = True success = True
break break
+11 -6
View File
@@ -55,20 +55,22 @@ class FutureMap:
self.future_buffer_len = self.future_limit + 2 * max_running_requests self.future_buffer_len = self.future_limit + 2 * max_running_requests
self.device = device self.device = device
self.spec_algo = spec_algo self.spec_algo = spec_algo
self.buf_initialized = False
if self.spec_algo.is_none(): if self.spec_algo.is_none():
# For non-speculative decoding, we only need to store the token ids.
self.buf_initialized = True
self.token_ids_buf = torch.empty( self.token_ids_buf = torch.empty(
(self.future_buffer_len,), dtype=torch.int64, device=self.device (self.future_buffer_len,), dtype=torch.int64, device=self.device
) )
else:
# For speculative decoding, we lazily initialize the buffers
# This is to make the shape derivation easier.
self.buf_initialized = False
def _lazy_init_buf(self, draft_input: EagleDraftInput): def _lazy_init_buf(self, draft_input: EagleDraftInput):
if self.buf_initialized or not self.spec_algo.is_eagle():
return
self.buf_initialized = True self.buf_initialized = True
# get the template for each tensor # Get a reference for each tensor
topk_p0 = draft_input.topk_p[0] topk_p0 = draft_input.topk_p[0]
topk_index0 = draft_input.topk_index[0] topk_index0 = draft_input.topk_index[0]
hidden_states0 = draft_input.hidden_states[0] hidden_states0 = draft_input.hidden_states[0]
@@ -147,10 +149,13 @@ class FutureMap:
self, future_indices: FutureIndices, draft_input: EagleDraftInput self, future_indices: FutureIndices, draft_input: EagleDraftInput
): ):
intv = future_indices.interval intv = future_indices.interval
# idle indices do not need store info
if self.is_empty_slice(intv): if self.is_empty_slice(intv):
# idle indices in dp attention do not need store info
return return
if not self.buf_initialized:
self._lazy_init_buf(draft_input) self._lazy_init_buf(draft_input)
self.topk_p_buf[intv] = draft_input.topk_p self.topk_p_buf[intv] = draft_input.topk_p
self.topk_index_buf[intv] = draft_input.topk_index self.topk_index_buf[intv] = draft_input.topk_index
self.hidden_states_buf[intv] = draft_input.hidden_states self.hidden_states_buf[intv] = draft_input.hidden_states
+5 -5
View File
@@ -1819,15 +1819,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
) )
@property @property
def is_v2_eagle(self): def is_eagle_v2(self):
# FIXME: finally deprecate is_v2_eagle # FIXME: finally deprecate is_eagle_v2
return self.enable_overlap and self.spec_algorithm.is_eagle() return self.enable_overlap and self.spec_algorithm.is_eagle()
def prepare_for_decode(self): def prepare_for_decode(self):
self.forward_mode = ForwardMode.DECODE self.forward_mode = ForwardMode.DECODE
bs = len(self.reqs) bs = len(self.reqs)
if self.is_v2_eagle: if self.is_eagle_v2:
# TODO(spec-v2): all v2 spec should go through this path # TODO(spec-v2): all v2 spec should go through this path
draft_input: EagleDraftInput = self.spec_info draft_input: EagleDraftInput = self.spec_info
draft_input.prepare_for_decode(self) draft_input.prepare_for_decode(self)
@@ -1907,7 +1907,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
) )
def maybe_wait_verify_done(self): def maybe_wait_verify_done(self):
if self.is_v2_eagle: if self.is_eagle_v2:
draft_input: EagleDraftInput = self.spec_info draft_input: EagleDraftInput = self.spec_info
if draft_input.verify_done is not None: if draft_input.verify_done is not None:
draft_input.verify_done.synchronize() draft_input.verify_done.synchronize()
@@ -1980,7 +1980,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# NOTE: spec_info filtered before batch filtering only happens in: # NOTE: spec_info filtered before batch filtering only happens in:
# - Spec v1's verify phase # - Spec v1's verify phase
# - Only for decode batch (running_batch) # - Only for decode batch (running_batch)
has_been_filtered = v1_spec_info_filtered and not self.is_v2_eagle has_been_filtered = v1_spec_info_filtered and not self.is_eagle_v2
if self.spec_info: if self.spec_info:
self.spec_info.filter_batch( self.spec_info.filter_batch(
+153 -162
View File
@@ -296,6 +296,9 @@ class Scheduler(
# Init diffusion LLM config # Init diffusion LLM config
self.dllm_config = DllmConfig.from_server_args(server_args) self.dllm_config = DllmConfig.from_server_args(server_args)
# Init metrics stats
self.init_metrics(tp_rank, pp_rank, dp_rank)
# Init inter-process communication # Init inter-process communication
self.init_sockets(server_args, port_args) self.init_sockets(server_args, port_args)
@@ -430,9 +433,6 @@ class Scheduler(
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB" f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB"
) )
# Init metrics stats
self.init_metrics(tp_rank, pp_rank, dp_rank)
# Init cache using the existing memory pool # Init cache using the existing memory pool
self.init_cache_with_memory_pool() self.init_cache_with_memory_pool()
@@ -447,18 +447,11 @@ class Scheduler(
# The last forward batch # The last forward batch
self.last_batch: Optional[ScheduleBatch] = None self.last_batch: Optional[ScheduleBatch] = None
self.forward_ct = 0 self.forward_ct = 0
self.forward_ct_decode = 0
self.num_generated_tokens = 0
self.last_prefill_tokens = 0 self.last_prefill_tokens = 0
self.return_health_check_ct = 0 self.return_health_check_ct = 0
self.num_retracted_reqs: int = 0 self.num_retracted_reqs: int = 0
self.num_paused_reqs: int = 0 self.num_paused_reqs: int = 0
self.sessions: Dict[str, Session] = {} self.sessions: Dict[str, Session] = {}
self.default_stream: CudaStream = torch.get_device_module(
self.device
).current_stream()
if self.device == "cpu":
self.default_stream.synchronize = lambda: None # No-op for CPU
self.forward_sleep_time = None self.forward_sleep_time = None
self._engine_paused = False self._engine_paused = False
@@ -696,7 +689,7 @@ class Scheduler(
self.send_to_tokenizer = SenderWrapper(None) self.send_to_tokenizer = SenderWrapper(None)
self.send_to_detokenizer = SenderWrapper(None) self.send_to_detokenizer = SenderWrapper(None)
if self.current_scheduler_metrics_enabled(): if self.current_scheduler_metrics_enabled:
self.send_metrics_from_scheduler = get_zmq_socket( self.send_metrics_from_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.metrics_ipc_name, False context, zmq.PUSH, port_args.metrics_ipc_name, False
) )
@@ -974,20 +967,22 @@ class Scheduler(
self.disagg_prefill_inflight_queue: List[Req] = [] self.disagg_prefill_inflight_queue: List[Req] = []
def init_overlap(self): def init_overlap(self):
self.future_map = None self.device_module = torch.get_device_module(self.device)
if not self.enable_overlap and self.pp_size == 1: self.default_stream: CudaStream = self.device_module.current_stream()
return if self.device == "cpu":
self.default_stream.synchronize = lambda: None # No-op for CPU
self.forward_stream: CudaStream = torch.get_device_module(self.device).Stream() self.forward_stream: CudaStream = self.device_module.Stream()
self.forward_stream_ctx: CudaStreamContext = torch.get_device_module( self.forward_stream_ctx: CudaStreamContext = self.device_module.stream(
self.device self.forward_stream
).stream(self.forward_stream) )
self.copy_stream: CudaStream = torch.get_device_module(self.device).Stream() self.copy_stream: CudaStream = self.device_module.Stream()
self.copy_stream_ctx: CudaStreamContext = torch.get_device_module( self.copy_stream_ctx: CudaStreamContext = self.device_module.stream(
self.device self.copy_stream
).stream(self.copy_stream) )
if not self.enable_overlap: if not self.enable_overlap:
self.future_map = None
return return
self.future_map = FutureMap( self.future_map = FutureMap(
@@ -1000,15 +995,6 @@ class Scheduler(
self.batch_record_buf = [None] * 2 self.batch_record_buf = [None] * 2
self.batch_record_ct = 0 self.batch_record_ct = 0
def record_batch_in_overlap(self, model_worker_batch: ModelWorkerBatch):
# FIXME(lsyin): hacky way to keep a reference to avoid GPU tensors being freed by torch GC
# NOTE: More Reliable: record all tensors into the forward stream
# NOTE: - for all future tensors, we shall always read from future map
# - for all non-future tensors (produced only by schedule stream),
# we shall keep its reference not being release during all the forwarding pass
self.batch_record_ct = (self.batch_record_ct + 1) % 2
self.batch_record_buf[self.batch_record_ct] = model_worker_batch
def init_moe_config(self): def init_moe_config(self):
if hasattr(self.model_config.hf_config, "num_experts_per_tok"): if hasattr(self.model_config.hf_config, "num_experts_per_tok"):
initialize_moe_config(self.server_args) initialize_moe_config(self.server_args)
@@ -1023,15 +1009,17 @@ class Scheduler(
def event_loop_normal(self): def event_loop_normal(self):
"""A normal scheduler loop.""" """A normal scheduler loop."""
while True: while True:
# Receive requests
recv_reqs = self.recv_requests() recv_reqs = self.recv_requests()
self.process_input_requests(recv_reqs) self.process_input_requests(recv_reqs)
if self._engine_paused: if self._engine_paused:
continue continue
# Get the next batch to run
batch = self.get_next_batch_to_run() batch = self.get_next_batch_to_run()
self.cur_batch = batch self.cur_batch = batch
# Launch the current batch
if batch: if batch:
result = self.run_batch(batch) result = self.run_batch(batch)
self.process_batch_result(batch, result) self.process_batch_result(batch, result)
@@ -1039,8 +1027,8 @@ class Scheduler(
# When the server is idle, do self-check and re-init some states # When the server is idle, do self-check and re-init some states
self.self_check_during_idle() self.self_check_during_idle()
# Update the last batch
self.last_batch = batch self.last_batch = batch
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get(): if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
self.self_check_during_busy() self.self_check_during_busy()
@@ -1048,9 +1036,6 @@ class Scheduler(
def event_loop_overlap(self): def event_loop_overlap(self):
"""A scheduler loop that overlaps the CPU processing and GPU computation.""" """A scheduler loop that overlaps the CPU processing and GPU computation."""
self.result_queue: Deque[Tuple[ScheduleBatch, GenerationBatchResult]] = deque() self.result_queue: Deque[Tuple[ScheduleBatch, GenerationBatchResult]] = deque()
disable_consecutive_prefill_overlap = (
envs.SGLANG_DISABLE_CONSECUTIVE_PREFILL_OVERLAP.get()
)
def pop_and_process(): def pop_and_process():
# Process the results of the last batch # Process the results of the last batch
@@ -1058,52 +1043,68 @@ class Scheduler(
self.process_batch_result(tmp_batch, tmp_result) self.process_batch_result(tmp_batch, tmp_result)
while True: while True:
# Receive requests
recv_reqs = self.recv_requests() recv_reqs = self.recv_requests()
self.process_input_requests(recv_reqs) self.process_input_requests(recv_reqs)
if self._engine_paused: if self._engine_paused:
continue continue
# Get the next batch to run
batch = self.get_next_batch_to_run() batch = self.get_next_batch_to_run()
self.cur_batch = batch self.cur_batch = batch
disable_overlap_for_batch = self.is_disable_overlap_for_batch(batch)
# If we do not need to overlap the current batch with the last batch,
# we can process the last batch immediately.
if disable_overlap_for_batch:
pop_and_process()
# Launch the current batch
batch_result = None
if batch:
batch_result = self.run_batch(batch)
self.result_queue.append((batch.copy(), batch_result))
# Process the last batch
if self.last_batch:
if not disable_overlap_for_batch:
pop_and_process()
elif batch is None:
# When the server is idle, do self-check and re-init some states
self.self_check_during_idle()
# Run sample of the current batch
# It depends on the result of the last batch (e.g., grammar), so we run it after the last batch is processed.
self.launch_batch_sample_if_needed(batch_result)
# Update the last batch
self.last_batch = batch
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
self.self_check_during_busy()
def is_disable_overlap_for_batch(self, batch: ScheduleBatch) -> bool:
# For two consecutive prefill batches, we disable overlap to improve the TTFT of the first batch.
# This might slightly hurt the throughput, so we use an environment variable to control it.
disable_overlap_for_batch = ( disable_overlap_for_batch = (
disable_consecutive_prefill_overlap envs.SGLANG_DISABLE_CONSECUTIVE_PREFILL_OVERLAP.get()
and batch and batch
and batch.forward_mode.is_extend() and batch.forward_mode.is_extend()
and self.last_batch and self.last_batch
and self.last_batch.forward_mode.is_extend() and self.last_batch.forward_mode.is_extend()
) )
# FIXME(lsyin): remove this grammar sync # We do not support overlap + spec + grammar yet,
# so we need to turn off overlap for this batch.
# TODO(lsyin): support overlap + spec + grammar
need_grammar_sync = ( need_grammar_sync = (
batch is not None batch
and batch.forward_mode.is_decode() and batch.is_eagle_v2
and batch.has_grammar and batch.has_grammar
and batch.is_v2_eagle and batch.forward_mode.is_decode()
and len(self.result_queue) > 0 and len(self.result_queue) > 0
) )
if disable_overlap_for_batch or need_grammar_sync: return disable_overlap_for_batch or need_grammar_sync
pop_and_process()
batch_result = None
if batch:
batch_result = self.run_batch(batch)
self.result_queue.append((batch.copy(), batch_result))
if self.last_batch:
if not disable_overlap_for_batch and not need_grammar_sync:
pop_and_process()
elif batch is None:
# When the server is idle, do self-check and re-init some states
self.self_check_during_idle()
self.launch_batch_sample_if_needed(batch_result)
self.last_batch = batch
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
self.self_check_during_busy()
def recv_limit_reached(self, num_recv_reqs: int) -> bool: def recv_limit_reached(self, num_recv_reqs: int) -> bool:
if self.max_recv_per_poll < 0: if self.max_recv_per_poll < 0:
@@ -1163,32 +1164,7 @@ class Scheduler(
if self.server_args.enable_dp_attention: if self.server_args.enable_dp_attention:
if self.attn_tp_rank == 0: if self.attn_tp_rank == 0:
work_reqs = [ work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
req
for req in recv_reqs
if isinstance(
req,
(
TokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
BatchTokenizedEmbeddingReqInput,
),
)
]
control_reqs = [
req
for req in recv_reqs
if not isinstance(
req,
(
TokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
BatchTokenizedEmbeddingReqInput,
),
)
]
else: else:
work_reqs = None work_reqs = None
control_reqs = None control_reqs = None
@@ -1226,9 +1202,37 @@ class Scheduler(
return recv_reqs return recv_reqs
def process_input_requests(self, recv_reqs: List): def _split_work_and_control_reqs(self, recv_reqs: List):
work_reqs = [
req
for req in recv_reqs
if isinstance(
req,
(
TokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
BatchTokenizedEmbeddingReqInput,
),
)
]
control_reqs = [
req
for req in recv_reqs
if not isinstance(
req,
(
TokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
BatchTokenizedEmbeddingReqInput,
),
)
]
return work_reqs, control_reqs
# Process MM requests under E disaggregation def process_input_requests(self, recv_reqs: List):
# Process MM requests under EPD-disaggregation mode
if ( if (
self.server_args.language_only self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
@@ -1247,11 +1251,11 @@ class Scheduler(
output = self._request_dispatcher(recv_req) output = self._request_dispatcher(recv_req)
if output is not None: if output is not None:
if isinstance(output, RpcReqOutput): if not isinstance(output, RpcReqOutput):
self.send_to_tokenizer.send_output(output, recv_req)
else:
if self.recv_from_rpc is not None: if self.recv_from_rpc is not None:
self.recv_from_rpc.send_pyobj(output) self.recv_from_rpc.send_pyobj(output)
else:
self.send_to_tokenizer.send_output(output, recv_req)
def init_req_max_new_tokens(self, req): def init_req_max_new_tokens(self, req):
req.sampling_params.max_new_tokens = min( req.sampling_params.max_new_tokens = min(
@@ -1782,6 +1786,7 @@ class Scheduler(
): ):
# Decrease prefill idle as much as possible during high dp load. # Decrease prefill idle as much as possible during high dp load.
return None return None
# 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_queue:
self.move_ready_grammar_requests() self.move_ready_grammar_requests()
@@ -1882,9 +1887,9 @@ class Scheduler(
self.running_batch.batch_is_full = True self.running_batch.batch_is_full = True
if self.running_batch.batch_is_full: if self.running_batch.batch_is_full:
if not self.try_preemption: if not self.try_preemption or not adder.preempt_to_schedule(
break req, self.server_args
if not adder.preempt_to_schedule(req, self.server_args): ):
break break
if self.enable_hicache_storage: if self.enable_hicache_storage:
@@ -1928,6 +1933,7 @@ class Scheduler(
for req in adder.preempt_list: for req in adder.preempt_list:
self._add_request_to_queue(req) self._add_request_to_queue(req)
# Update chunked prefill
if adder.new_chunked_req is not None: if adder.new_chunked_req is not None:
assert self.chunked_req is None assert self.chunked_req is None
self.chunked_req = adder.new_chunked_req self.chunked_req = adder.new_chunked_req
@@ -1936,12 +1942,12 @@ class Scheduler(
self.chunked_req.is_chunked += 1 self.chunked_req.is_chunked += 1
# Print stats # Print stats
if self.current_scheduler_metrics_enabled(): if self.current_scheduler_metrics_enabled:
self.log_prefill_stats(adder, can_run_list, running_bs, 0) self.log_prefill_stats(adder, can_run_list, running_bs, 0)
# Record metrics
for req in can_run_list: for req in can_run_list:
if req.time_stats.forward_entry_time == 0: if req.time_stats.forward_entry_time == 0:
# Avoid update chunked request many times
req.time_stats.forward_entry_time = time.perf_counter() req.time_stats.forward_entry_time = time.perf_counter()
if self.enable_metrics: if self.enable_metrics:
self.metrics_collector.observe_queue_time( self.metrics_collector.observe_queue_time(
@@ -2034,11 +2040,14 @@ class Scheduler(
batch.prepare_for_decode() batch.prepare_for_decode()
return batch return batch
# placeholder for override def record_batch_in_overlap(self, model_worker_batch: ModelWorkerBatch):
def update_cache_from_scheduler( # FIXME(lsyin): hacky way to keep a reference to avoid GPU tensors being freed by torch GC
self, schedule_batch: ScheduleBatch, batch_result: GenerationBatchResult # NOTE: More Reliable: record all tensors into the forward stream
): # NOTE: - for all future tensors, we shall always read from future map
pass # - for all non-future tensors (produced only by schedule stream),
# we shall keep its reference not being release during all the forwarding pass
self.batch_record_ct = (self.batch_record_ct + 1) % 2
self.batch_record_buf[self.batch_record_ct] = model_worker_batch
def run_batch( def run_batch(
self, self,
@@ -2066,19 +2075,19 @@ class Scheduler(
# Run forward # Run forward
if self.is_generation: if self.is_generation:
batch_or_worker_batch = batch if self.spec_algorithm.is_none() or self.enable_overlap:
# In most cases, we use the model worker batch to run the forward.
if self.enable_overlap or self.spec_algorithm.is_none(): worker_batch_or_batch = batch.get_model_worker_batch()
# FIXME(lsyin): remove this if and finally unify the abstraction else:
batch_or_worker_batch = batch.get_model_worker_batch() # In speculative decoding v1 (non-overlap) case, we use the batch directly.
# TODO(lsyin): delete this branch after unifying the abstraction.
worker_batch_or_batch = batch
if self.enable_overlap: if self.enable_overlap:
# FIXME: remove this assert model_worker_batch = worker_batch_or_batch
assert isinstance(batch_or_worker_batch, ModelWorkerBatch)
model_worker_batch = batch_or_worker_batch
self.record_batch_in_overlap(model_worker_batch) self.record_batch_in_overlap(model_worker_batch)
# Sampling info will be modified during forward # Sampling info will be modified during forward, so we store a copy.
model_worker_batch.sampling_info = ( model_worker_batch.sampling_info = (
model_worker_batch.sampling_info.copy_for_forward() model_worker_batch.sampling_info.copy_for_forward()
) )
@@ -2094,9 +2103,7 @@ class Scheduler(
# here pp is not compatible with overlap # here pp is not compatible with overlap
) )
# FIXME(lsyin): maybe move this to forward_batch_generation # FIXME(lsyin): maybe move this to forward_batch_generation
batch_result.copy_done = torch.get_device_module( batch_result.copy_done = self.device_module.Event()
self.device
).Event()
if batch_result.delay_sample_func is None: if batch_result.delay_sample_func is None:
self.future_map.store_to_map(future_indices, batch_result) self.future_map.store_to_map(future_indices, batch_result)
batch_result.copy_to_cpu(return_logprob=batch.return_logprob) batch_result.copy_to_cpu(return_logprob=batch.return_logprob)
@@ -2106,7 +2113,7 @@ class Scheduler(
# FIXME(lsyin): move this assignment elsewhere # FIXME(lsyin): move this assignment elsewhere
future_indices_or_next_token_ids = -future_indices.indices future_indices_or_next_token_ids = -future_indices.indices
if batch.is_v2_eagle: if batch.is_eagle_v2:
# FIXME(lsyin): tmp code for eagle v2 # FIXME(lsyin): tmp code for eagle v2
# We only keep future indices for next draft input # We only keep future indices for next draft input
@@ -2131,7 +2138,7 @@ class Scheduler(
else {} else {}
) )
batch_result = self.model_worker.forward_batch_generation( batch_result = self.model_worker.forward_batch_generation(
batch_or_worker_batch, **kwargs worker_batch_or_batch, **kwargs
) )
future_indices_or_next_token_ids = batch_result.next_token_ids future_indices_or_next_token_ids = batch_result.next_token_ids
self.update_cache_from_scheduler(batch, batch_result) self.update_cache_from_scheduler(batch, batch_result)
@@ -2146,21 +2153,16 @@ class Scheduler(
# modified by overlap schedule. So we have to copy them here so that # modified by overlap schedule. So we have to copy them here so that
# we can use the correct values in output processing. # we can use the correct values in output processing.
if batch.return_logprob or self.spec_algorithm.is_eagle(): if batch.return_logprob or self.spec_algorithm.is_eagle():
extend_input_len_per_req = [req.extend_input_len for req in batch.reqs] batch_result.extend_input_len_per_req = [
else: req.extend_input_len for req in batch.reqs
extend_input_len_per_req = None ]
batch_result.extend_logprob_start_len_per_req = [
if batch.return_logprob:
extend_logprob_start_len_per_req = [
req.extend_logprob_start_len for req in batch.reqs req.extend_logprob_start_len for req in batch.reqs
] ]
else: else:
extend_logprob_start_len_per_req = None batch_result.extend_input_len_per_req = None
batch_result.extend_logprob_start_len_per_req = None
batch_result.extend_input_len_per_req = extend_input_len_per_req
batch_result.extend_logprob_start_len_per_req = (
extend_logprob_start_len_per_req
)
ret = batch_result ret = batch_result
else: # embedding or reward model else: # embedding or reward model
model_worker_batch = batch.get_model_worker_batch() model_worker_batch = batch.get_model_worker_batch()
@@ -2327,31 +2329,27 @@ class Scheduler(
self.cur_batch = None self.cur_batch = None
self.last_batch = None self.last_batch = None
self.tree_cache.reset() self.tree_cache.reset()
if self.grammar_backend:
self.grammar_backend.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_backend.reset()
self.reset_metrics()
if self.draft_worker: if self.draft_worker:
self.draft_worker.clear_cache_pool() self.draft_worker.clear_cache_pool()
self.num_generated_tokens = 0 # TODO: allow optional empty cache
self.forward_ct_decode = 0
self.spec_num_accepted_tokens = 0
self.spec_num_forward_ct = 0
self.spec_total_num_accepted_tokens = 0
self.spec_total_num_forward_ct = 0
torch.cuda.empty_cache() torch.cuda.empty_cache()
logger.info("Cache flushed successfully!") logger.info("Cache flushed successfully!")
if_success = True success = True
else: else:
logging.warning( logging.warning(
f"Cache not flushed because there are pending requests. " f"Cache not flushed because there are pending requests. "
f"#queue-req: {len(self.waiting_queue)}, " f"#queue-req: {len(self.waiting_queue)}, "
f"#running-req: {len(self.running_batch.reqs)}" f"#running-req: {len(self.running_batch.reqs)}"
) )
if_success = False success = False
return if_success return success
def get_internal_state(self, recv_req: GetInternalStateReq): def get_internal_state(self, recv_req: GetInternalStateReq):
ret = vars(get_global_server_args()) ret = vars(get_global_server_args())
@@ -2362,16 +2360,14 @@ class Scheduler(
self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 2 self.token_to_kv_pool_allocator.get_kvcache().mem_usage, 2
), ),
"token_capacity": int(self.max_total_num_tokens), "token_capacity": int(self.max_total_num_tokens),
"graph": round(self.tp_worker.model_runner.graph_mem_usage, 2),
} }
ret["memory_usage"]["graph"] = round(
self.tp_worker.model_runner.graph_mem_usage, 2
)
if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0: if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0:
ret["avg_spec_accept_length"] = ( ret["avg_spec_accept_length"] = (
self.spec_total_num_accepted_tokens / self.spec_total_num_forward_ct self.spec_total_num_accepted_tokens / self.spec_total_num_forward_ct
) )
if RECORD_STEP_TIME: if RECORD_STEP_TIME:
ret["step_time_dict"] = self.step_time_dict ret["step_time_dict"] = self.step_time_dict
@@ -2389,6 +2385,7 @@ class Scheduler(
"speculative_accept_threshold_acc", "speculative_accept_threshold_acc",
] ]
) )
if_success = True if_success = True
for k, v in server_args_dict.items(): for k, v in server_args_dict.items():
if k not in args_allow_update: if k not in args_allow_update:
@@ -2403,6 +2400,7 @@ class Scheduler(
) )
if_success = False if_success = False
break break
if if_success: if if_success:
if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0: if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0:
avg_spec_accept_length = ( avg_spec_accept_length = (
@@ -2633,19 +2631,6 @@ class Scheduler(
else: else:
del self.sessions[session_id] del self.sessions[session_id]
def get_print_prefix(self):
prefix = ""
if self.attn_dp_rank is not None:
prefix += f" DP{self.attn_dp_rank}"
if self.server_args.tp_size > 1:
prefix += f" TP{self.tp_rank}"
if self.pp_size > 1:
prefix += f" PP{self.pp_rank}"
return prefix
def current_scheduler_metrics_enabled(self):
return self.attn_tp_rank == 0 or self.enable_metrics_for_all_schedulers
def maybe_sleep_on_idle(self): def maybe_sleep_on_idle(self):
if self.idle_sleeper is not None: if self.idle_sleeper is not None:
self.idle_sleeper.maybe_sleep() self.idle_sleeper.maybe_sleep()
@@ -2656,6 +2641,12 @@ class Scheduler(
self.send_to_detokenizer.send_output(recv_req, recv_req) self.send_to_detokenizer.send_output(recv_req, recv_req)
return None return None
# placeholder for override
def update_cache_from_scheduler(
self, schedule_batch: ScheduleBatch, batch_result: GenerationBatchResult
):
pass
def get_remote_instance_transfer_engine_info(self): def get_remote_instance_transfer_engine_info(self):
return self.tp_worker.get_remote_instance_transfer_engine_info() return self.tp_worker.get_remote_instance_transfer_engine_info()
@@ -2791,6 +2782,8 @@ def run_scheduler_process(
) )
pipe_writer.send(result_dict) pipe_writer.send(result_dict)
# Dispatch to the appropriate event loop based on the disaggregation mode
disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
if disaggregation_mode == DisaggregationMode.NULL: if disaggregation_mode == DisaggregationMode.NULL:
if scheduler.enable_pdmux: if scheduler.enable_pdmux:
@@ -2802,20 +2795,18 @@ def run_scheduler_process(
else: else:
scheduler.event_loop_normal() scheduler.event_loop_normal()
elif disaggregation_mode == DisaggregationMode.PREFILL: elif disaggregation_mode == DisaggregationMode.PREFILL:
if scheduler.enable_overlap:
scheduler.event_loop_overlap_disagg_prefill()
else:
if server_args.pp_size > 1: if server_args.pp_size > 1:
scheduler.event_loop_pp_disagg_prefill() scheduler.event_loop_pp_disagg_prefill()
elif scheduler.enable_overlap:
scheduler.event_loop_overlap_disagg_prefill()
else: else:
scheduler.event_loop_normal_disagg_prefill() scheduler.event_loop_normal_disagg_prefill()
elif disaggregation_mode == DisaggregationMode.DECODE: elif disaggregation_mode == DisaggregationMode.DECODE:
if scheduler.enable_overlap:
scheduler.event_loop_overlap_disagg_decode()
else:
if server_args.pp_size > 1: if server_args.pp_size > 1:
scheduler.event_loop_pp_disagg_decode() scheduler.event_loop_pp_disagg_decode()
elif scheduler.enable_overlap:
scheduler.event_loop_overlap_disagg_decode()
else: else:
scheduler.event_loop_normal_disagg_decode() scheduler.event_loop_normal_disagg_decode()
@@ -39,9 +39,11 @@ class SchedulerMetricsMixin:
def init_metrics( def init_metrics(
self: Scheduler, tp_rank: int, pp_rank: int, dp_rank: Optional[int] self: Scheduler, tp_rank: int, pp_rank: int, dp_rank: Optional[int]
): ):
# Basic stats
self.forward_ct_decode = 0
self.num_generated_tokens = 0
self.last_decode_stats_tic = time.perf_counter() self.last_decode_stats_tic = time.perf_counter()
self.last_prefill_stats_tic = time.perf_counter() self.last_prefill_stats_tic = time.perf_counter()
self.last_gen_throughput: float = 0.0 self.last_gen_throughput: float = 0.0
self.last_input_throughput: float = 0.0 self.last_input_throughput: float = 0.0
self.step_time_dict = defaultdict(list) # Dict[batch size -> step time] self.step_time_dict = defaultdict(list) # Dict[batch size -> step time]
@@ -52,6 +54,8 @@ class SchedulerMetricsMixin:
# The total number of accepted tokens and forward ct for the whole server lifetime # The total number of accepted tokens and forward ct for the whole server lifetime
self.spec_total_num_accepted_tokens = 0 self.spec_total_num_accepted_tokens = 0
self.spec_total_num_forward_ct = 0 self.spec_total_num_forward_ct = 0
# For PD disaggregation
self.kv_transfer_speed_gb_s: float = 0.0 self.kv_transfer_speed_gb_s: float = 0.0
self.kv_transfer_latency_ms: float = 0.0 self.kv_transfer_latency_ms: float = 0.0
self.kv_transfer_bootstrap_ms: float = 0.0 self.kv_transfer_bootstrap_ms: float = 0.0
@@ -59,6 +63,11 @@ class SchedulerMetricsMixin:
self.stats = SchedulerStats() self.stats = SchedulerStats()
# Metrics
self.current_scheduler_metrics_enabled = (
self.attn_tp_rank == 0 or self.enable_metrics_for_all_schedulers
)
if self.enable_metrics: if self.enable_metrics:
engine_type = "unified" engine_type = "unified"
labels = { labels = {
@@ -82,6 +91,14 @@ class SchedulerMetricsMixin:
self.spec_num_forward_ct += bs self.spec_num_forward_ct += bs
self.num_generated_tokens += num_accepted_tokens self.num_generated_tokens += num_accepted_tokens
def reset_metrics(self):
self.forward_ct_decode = 0
self.num_generated_tokens = 0
self.spec_num_accepted_tokens = 0
self.spec_num_forward_ct = 0
self.spec_total_num_accepted_tokens = 0
self.spec_total_num_forward_ct = 0
def log_prefill_stats( def log_prefill_stats(
self: Scheduler, self: Scheduler,
adder: PrefillAdder, adder: PrefillAdder,
@@ -331,7 +331,7 @@ class SchedulerOutputProcessorMixin:
next_token_ids = next_token_ids.tolist() next_token_ids = next_token_ids.tolist()
if batch.return_logprob: if batch.return_logprob:
next_token_logprobs = logits_output.next_token_logprobs.tolist() next_token_logprobs = logits_output.next_token_logprobs.tolist()
elif batch.is_v2_eagle: elif batch.is_eagle_v2:
next_token_ids = self._resolve_spec_overlap_token_ids(result, batch) next_token_ids = self._resolve_spec_overlap_token_ids(result, batch)
self.num_generated_tokens += len(batch.reqs) self.num_generated_tokens += len(batch.reqs)
@@ -356,7 +356,7 @@ class SchedulerOutputProcessorMixin:
new_accepted_len = 1 new_accepted_len = 1
if batch.spec_algorithm.is_none(): if batch.spec_algorithm.is_none():
req.output_ids.append(next_token_id) req.output_ids.append(next_token_id)
elif batch.is_v2_eagle: elif batch.is_eagle_v2:
# Only v2 eagle's output_ids are updated here. # Only v2 eagle's output_ids are updated here.
req.output_ids.extend(next_token_id) req.output_ids.extend(next_token_id)
new_accepted_len = len(next_token_id) new_accepted_len = len(next_token_id)
@@ -406,7 +406,7 @@ class SchedulerOutputProcessorMixin:
if batch.spec_algorithm.is_none(): if batch.spec_algorithm.is_none():
# Normal decode: single token # Normal decode: single token
req.grammar.accept_token(next_token_id) req.grammar.accept_token(next_token_id)
elif batch.is_v2_eagle: elif batch.is_eagle_v2:
# Speculative decode: next_token_id is a list of accepted tokens # Speculative decode: next_token_id is a list of accepted tokens
for token_id in next_token_id: for token_id in next_token_id:
req.grammar.accept_token(token_id) req.grammar.accept_token(token_id)
@@ -424,7 +424,7 @@ class SchedulerOutputProcessorMixin:
self.forward_ct_decode = (self.forward_ct_decode + 1) % (1 << 30) self.forward_ct_decode = (self.forward_ct_decode + 1) % (1 << 30)
if ( if (
self.current_scheduler_metrics_enabled() self.current_scheduler_metrics_enabled
and self.forward_ct_decode % self.server_args.decode_log_interval == 0 and self.forward_ct_decode % self.server_args.decode_log_interval == 0
): ):
self.log_decode_stats(can_run_cuda_graph, running_batch=batch) self.log_decode_stats(can_run_cuda_graph, running_batch=batch)
@@ -246,7 +246,7 @@ class SchedulerRuntimeCheckerMixin:
if ( if (
self.enable_metrics self.enable_metrics
and self.current_scheduler_metrics_enabled() and self.current_scheduler_metrics_enabled
and time.perf_counter() > self.metrics_collector.last_log_time + 30 and time.perf_counter() > self.metrics_collector.last_log_time + 30
): ):
# During idle time, also collect metrics every 30 seconds. # During idle time, also collect metrics every 30 seconds.
+8 -6
View File
@@ -383,6 +383,7 @@ class TpModelWorker(BaseTpWorker):
# FIXME(lsyin): maybe remove skip_attn_backend_init in forward_batch_generation, # FIXME(lsyin): maybe remove skip_attn_backend_init in forward_batch_generation,
# which requires preparing replay to always be in this function # which requires preparing replay to always be in this function
# Get forward batch from model worker batch
if model_worker_batch is not None: if model_worker_batch is not None:
# update the consumer index of hicache to the running batch # update the consumer index of hicache to the running batch
self.set_hicache_consumer(model_worker_batch.hicache_consumer_index) self.set_hicache_consumer(model_worker_batch.hicache_consumer_index)
@@ -392,10 +393,10 @@ class TpModelWorker(BaseTpWorker):
# FIXME(lsyin): unify the interface of forward_batch # FIXME(lsyin): unify the interface of forward_batch
assert forward_batch is not None assert forward_batch is not None
if self.pp_group.is_last_rank:
if self.is_dllm(): if self.is_dllm():
return self._forward_batch_generation_dllm(forward_batch) return self._forward_batch_generation_dllm(forward_batch)
if self.pp_group.is_last_rank:
logits_output, can_run_cuda_graph = self.model_runner.forward( logits_output, can_run_cuda_graph = self.model_runner.forward(
forward_batch, forward_batch,
pp_proxy_tensors=pp_proxy_tensors, pp_proxy_tensors=pp_proxy_tensors,
@@ -425,7 +426,12 @@ class TpModelWorker(BaseTpWorker):
batch_result.delay_sample_func = sample_batch_func batch_result.delay_sample_func = sample_batch_func
return batch_result return batch_result
if model_worker_batch.is_prefill_only: if not model_worker_batch.is_prefill_only:
# For normal requests, sample the next token ids.
batch_result.next_token_ids = self.model_runner.sample(
logits_output, forward_batch
)
else:
# For prefill-only requests, create dummy token IDs on CPU # For prefill-only requests, create dummy token IDs on CPU
# The size should match the batch size (number of sequences), not total tokens # The size should match the batch size (number of sequences), not total tokens
batch_result.next_token_ids = torch.zeros( batch_result.next_token_ids = torch.zeros(
@@ -441,10 +447,6 @@ class TpModelWorker(BaseTpWorker):
self.model_runner.compute_logprobs_only( self.model_runner.compute_logprobs_only(
logits_output, model_worker_batch logits_output, model_worker_batch
) )
else:
batch_result.next_token_ids = self.model_runner.sample(
logits_output, forward_batch
)
return batch_result return batch_result
else: else:
+1
View File
@@ -1966,6 +1966,7 @@ class ServerArgs:
self.disable_overlap_schedule = True self.disable_overlap_schedule = True
logger.warning( logger.warning(
"Overlap scheduler is disabled because of using eagle3 or standalone speculative decoding." "Overlap scheduler is disabled because of using eagle3 or standalone speculative decoding."
"You can set env SGLANG_ENABLE_SPEC_V2=True to enable the experimental overlap scheduler."
) )
if self.enable_mixed_chunk: if self.enable_mixed_chunk:
+2 -12
View File
@@ -43,6 +43,7 @@ fi
# Install protoc for router build (gRPC protobuf compilation) # Install protoc for router build (gRPC protobuf compilation)
if [ "${INSTALL_PROTOC:-0}" = "1" ]; then if [ "${INSTALL_PROTOC:-0}" = "1" ]; then
# TODO: move this to a separate script
echo "Installing protoc..." echo "Installing protoc..."
if command -v apt-get &> /dev/null; then if command -v apt-get &> /dev/null; then
# Ubuntu/Debian # Ubuntu/Debian
@@ -108,6 +109,7 @@ $PIP_CMD install -e "python[${EXTRAS}]" --extra-index-url https://download.pytor
# Install router for pd-disagg test # Install router for pd-disagg test
$PIP_CMD install sglang-router $PIP_INSTALL_SUFFIX $PIP_CMD install sglang-router $PIP_INSTALL_SUFFIX
# Remove flash_attn folder to avoid conflicts
PYTHON_LIB_PATH=$(python3 -c "import site; print(site.getsitepackages()[0])") PYTHON_LIB_PATH=$(python3 -c "import site; print(site.getsitepackages()[0])")
FLASH_ATTN_PATH="${PYTHON_LIB_PATH}/flash_attn" FLASH_ATTN_PATH="${PYTHON_LIB_PATH}/flash_attn"
@@ -215,15 +217,3 @@ python3 -c "import torch; print(torch.version.cuda)"
# Prepare the CI runner (cleanup HuggingFace cache, etc.) # Prepare the CI runner (cleanup HuggingFace cache, etc.)
bash "${SCRIPT_DIR}/prepare_runner.sh" bash "${SCRIPT_DIR}/prepare_runner.sh"
# Remove flash_attn folder to avoid conflicts with sgl-kernel
PYTHON_LIB_PATH=$(python3 -c "import site; print(site.getsitepackages()[0])")
FLASH_ATTN_PATH="${PYTHON_LIB_PATH}/flash_attn"
if [ -d "$FLASH_ATTN_PATH" ]; then
echo "Directory $FLASH_ATTN_PATH exists. Removing..."
rm -rf "$FLASH_ATTN_PATH"
echo "error: this should not happen"
else
echo "Directory $FLASH_ATTN_PATH does not exist."
fi