[style] Extract init-static values in scheduler hot path (#30707)

This commit is contained in:
Liangsheng Yin
2026-07-09 19:33:02 -07:00
committed by GitHub
parent 073b36853f
commit b5e75b9423
4 changed files with 24 additions and 23 deletions
+10 -7
View File
@@ -2092,6 +2092,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def prepare_for_extend(self): def prepare_for_extend(self):
self.forward_mode = ForwardMode.EXTEND self.forward_mode = ForwardMode.EXTEND
server_args = get_server_args()
if self.is_dllm(): if self.is_dllm():
# For DLLM, we use a separate forward mode # For DLLM, we use a separate forward mode
@@ -2225,7 +2226,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
req.already_computed = seq_len req.already_computed = seq_len
req.is_retracted = False req.is_retracted = False
if get_server_args().enable_mamba_extra_buffer(): if server_args.enable_mamba_extra_buffer():
track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req) track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req)
mamba_track_mask_cpu.append(track_entry.track_mask) mamba_track_mask_cpu.append(track_entry.track_mask)
mamba_track_indices_cpu.append(track_entry.track_index) mamba_track_indices_cpu.append(track_entry.track_index)
@@ -2330,7 +2331,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.extend_logprob_start_lens = extend_logprob_start_lens self.extend_logprob_start_lens = extend_logprob_start_lens
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids self.extend_input_logprob_token_ids = extend_input_logprob_token_ids
if get_server_args().enable_mamba_extra_buffer(): if server_args.enable_mamba_extra_buffer():
self.mamba_track_indices = torch.tensor( self.mamba_track_indices = torch.tensor(
mamba_track_indices_cpu, mamba_track_indices_cpu,
dtype=torch.int64, dtype=torch.int64,
@@ -2364,7 +2365,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self, self,
req: Req, req: Req,
) -> _MambaRadixCacheV2TrackEntry: ) -> _MambaRadixCacheV2TrackEntry:
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size server_args = get_server_args()
mamba_cache_chunk_size = server_args.mamba_cache_chunk_size
def _force_track_h(i: int) -> int: def _force_track_h(i: int) -> int:
assert i % mamba_cache_chunk_size == 0 assert i % mamba_cache_chunk_size == 0
@@ -2415,7 +2417,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# In lazy mode, skip the swap — the second ping-pong slot is not # In lazy mode, skip the swap — the second ping-pong slot is not
# allocated yet; it will be allocated on demand at the track boundary # allocated yet; it will be allocated on demand at the track boundary
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode. # in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
if not get_server_args().enable_mamba_extra_buffer_lazy(): if not server_args.enable_mamba_extra_buffer_lazy():
req.mamba_next_track_idx = ( req.mamba_next_track_idx = (
self.req_to_token_pool.get_mamba_ping_pong_other_idx( self.req_to_token_pool.get_mamba_ping_pong_other_idx(
req.mamba_next_track_idx req.mamba_next_track_idx
@@ -2738,6 +2740,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def prepare_for_decode(self): def prepare_for_decode(self):
self.forward_mode = ForwardMode.DECODE self.forward_mode = ForwardMode.DECODE
server_args = get_server_args()
# Decode embeds the last output token via embed_tokens; clear the stale # Decode embeds the last output token via embed_tokens; clear the stale
# prefill-time tensor so it doesn't leak into ForwardBatch. # prefill-time tensor so it doesn't leak into ForwardBatch.
self.input_embeds = None self.input_embeds = None
@@ -2794,15 +2797,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.req_pool_indices_cpu, self.req_pool_indices_cpu,
) )
if get_server_args().enable_mamba_extra_buffer(): if server_args.enable_mamba_extra_buffer():
mamba_track_interval = get_server_args().mamba_track_interval mamba_track_interval = server_args.mamba_track_interval
if len(self.reqs) == 0: if len(self.reqs) == 0:
self.mamba_track_indices = torch.empty( self.mamba_track_indices = torch.empty(
(0,), dtype=torch.int64, device=self.device (0,), dtype=torch.int64, device=self.device
) )
else: else:
if get_server_args().enable_mamba_extra_buffer_lazy(): if server_args.enable_mamba_extra_buffer_lazy():
self.mamba_lazy_prealloc_at_boundary(mamba_track_interval) self.mamba_lazy_prealloc_at_boundary(mamba_track_interval)
set_mamba_track_indices_from_reqs(self) set_mamba_track_indices_from_reqs(self)
+8 -9
View File
@@ -357,6 +357,8 @@ class Scheduler(
) )
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get() self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
self.enable_hisparse = server_args.enable_hisparse self.enable_hisparse = server_args.enable_hisparse
self.enable_dp_attention = server_args.enable_dp_attention
self.enable_unified_memory = server_args.enable_unified_memory
# Distributed rank info # Distributed rank info
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
@@ -470,7 +472,7 @@ class Scheduler(
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tp_group=( tp_group=(
self.attn_tp_cpu_group self.attn_tp_cpu_group
if self.server_args.enable_dp_attention if self.enable_dp_attention
else self.tp_cpu_group else self.tp_cpu_group
), ),
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
@@ -911,9 +913,7 @@ class Scheduler(
# Use the CPU (gloo) group to broadcast VLM Python objects and avoid CUDA # Use the CPU (gloo) group to broadcast VLM Python objects and avoid CUDA
# stream/device coupling (#11910). # stream/device coupling (#11910).
self.dp_tp_group = ( self.dp_tp_group = (
self.attn_tp_group self.attn_tp_group if self.enable_dp_attention else self.tp_group
if self.server_args.enable_dp_attention
else self.tp_group
) )
self.dp_tp_cpu_group = self.dp_tp_group.cpu_group self.dp_tp_cpu_group = self.dp_tp_group.cpu_group
@@ -1604,7 +1604,7 @@ class Scheduler(
# Opportunistic flush at the disable_overlap sync boundary: # Opportunistic flush at the disable_overlap sync boundary:
# forward_stream is idle (prev forward drained, next not launched), # forward_stream is idle (prev forward drained, next not launched),
# so `_flush`'s non-urgent guard compacts freely. Sync-free, best-effort. # so `_flush`'s non-urgent guard compacts freely. Sync-free, best-effort.
if self.server_args.enable_unified_memory: if self.enable_unified_memory:
try: try:
self.token_to_kv_pool_allocator.flush_opportunistic() self.token_to_kv_pool_allocator.flush_opportunistic()
except Exception: except Exception:
@@ -3272,7 +3272,7 @@ class Scheduler(
self.batch_record_buf[self.batch_record_ct].extend( self.batch_record_buf[self.batch_record_ct].extend(
batch_result.extra_keep_alive_refs batch_result.extra_keep_alive_refs
) )
if self.server_args.enable_unified_memory: if self.enable_unified_memory:
# Record a `forward_done` event after the forward (before # Record a `forward_done` event after the forward (before
# copy_to_cpu); lazy-compaction `_flush` gates src reuse on # copy_to_cpu); lazy-compaction `_flush` gates src reuse on
# it. Only the unified pool's allocator exposes these hooks. # it. Only the unified pool's allocator exposes these hooks.
@@ -3401,8 +3401,7 @@ class Scheduler(
def _maybe_report_active_ranks(self) -> None: def _maybe_report_active_ranks(self) -> None:
if not ( if not (
self.server_args.enable_dp_attention self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
and self.server_args.elastic_ep_backend is not None
): ):
return return
# Get the tensors indicating rank activeness # Get the tensors indicating rank activeness
@@ -3546,7 +3545,7 @@ class Scheduler(
if not self.is_fully_idle(): if not self.is_fully_idle():
return return
if self.server_args.enable_unified_memory: if self.enable_unified_memory:
try: try:
self.token_to_kv_pool_allocator.flush_opportunistic() self.token_to_kv_pool_allocator.flush_opportunistic()
except Exception: except Exception:
@@ -145,6 +145,7 @@ class SchedulerMetricsReporter:
self.kv_transfer_latency_ms: float = 0.0 self.kv_transfer_latency_ms: float = 0.0
self.enable_mfu_metrics = False self.enable_mfu_metrics = False
self.decode_log_interval = self.scheduler.server_args.decode_log_interval
if self.enable_metrics: if self.enable_metrics:
self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics
@@ -696,7 +697,7 @@ class SchedulerMetricsReporter:
x.maybe_dump(batch, self.scheduler.waiting_queue) x.maybe_dump(batch, self.scheduler.waiting_queue)
# Periodic work: log + heavy metrics at decode_log_interval # Periodic work: log + heavy metrics at decode_log_interval
if self.forward_ct_decode % self.scheduler.server_args.decode_log_interval != 0: if self.forward_ct_decode % self.decode_log_interval != 0:
return return
if ( if (
not self.is_stats_logging_rank not self.is_stats_logging_rank
@@ -716,7 +717,7 @@ class SchedulerMetricsReporter:
if RECORD_STEP_TIME: if RECORD_STEP_TIME:
self.step_time_dict[num_running_reqs].append( self.step_time_dict[num_running_reqs].append(
gap_latency / self.scheduler.server_args.decode_log_interval gap_latency / self.decode_log_interval
) )
batch_iter = ( batch_iter = (
@@ -792,7 +793,7 @@ class SchedulerMetricsReporter:
) )
msg += self._decode_sol_suffix( msg += self._decode_sol_suffix(
batch, batch,
gap_latency / max(1, self.scheduler.server_args.decode_log_interval), gap_latency / max(1, self.decode_log_interval),
) )
self._mfu_log_flops = 0.0 self._mfu_log_flops = 0.0
self._mfu_log_read_bytes = 0.0 self._mfu_log_read_bytes = 0.0
@@ -1033,10 +1034,7 @@ class SchedulerMetricsReporter:
self._device_timer_window_gpu_time / cpu_time * 100, 100 self._device_timer_window_gpu_time / cpu_time * 100, 100
) )
self._device_timer_window_batch_count += 1 self._device_timer_window_batch_count += 1
if ( if self._device_timer_window_batch_count >= self.decode_log_interval:
self._device_timer_window_batch_count
>= self.scheduler.server_args.decode_log_interval
):
self._device_timer_window_batch_count = 0 self._device_timer_window_batch_count = 0
def reset_device_timer_window(self): def reset_device_timer_window(self):
@@ -87,6 +87,7 @@ class _DummyPublisherThread:
def _fake_server_args(**fields): def _fake_server_args(**fields):
"""server_args stand-in: carries fields and the override() entry point.""" """server_args stand-in: carries fields and the override() entry point."""
fields.setdefault("decode_log_interval", 40)
ns = types.SimpleNamespace(**fields) ns = types.SimpleNamespace(**fields)
def _override(source, **updates): def _override(source, **updates):