[style] Extract init-static values in scheduler hot path (#30707)
This commit is contained in:
@@ -2092,6 +2092,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
|
||||
def prepare_for_extend(self):
|
||||
self.forward_mode = ForwardMode.EXTEND
|
||||
server_args = get_server_args()
|
||||
|
||||
if self.is_dllm():
|
||||
# For DLLM, we use a separate forward mode
|
||||
@@ -2225,7 +2226,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
req.already_computed = seq_len
|
||||
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)
|
||||
mamba_track_mask_cpu.append(track_entry.track_mask)
|
||||
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_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(
|
||||
mamba_track_indices_cpu,
|
||||
dtype=torch.int64,
|
||||
@@ -2364,7 +2365,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self,
|
||||
req: Req,
|
||||
) -> _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:
|
||||
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
|
||||
# allocated yet; it will be allocated on demand at the track boundary
|
||||
# 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 = (
|
||||
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
||||
req.mamba_next_track_idx
|
||||
@@ -2738,6 +2740,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
|
||||
def prepare_for_decode(self):
|
||||
self.forward_mode = ForwardMode.DECODE
|
||||
server_args = get_server_args()
|
||||
# Decode embeds the last output token via embed_tokens; clear the stale
|
||||
# prefill-time tensor so it doesn't leak into ForwardBatch.
|
||||
self.input_embeds = None
|
||||
@@ -2794,15 +2797,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.req_pool_indices_cpu,
|
||||
)
|
||||
|
||||
if get_server_args().enable_mamba_extra_buffer():
|
||||
mamba_track_interval = get_server_args().mamba_track_interval
|
||||
if server_args.enable_mamba_extra_buffer():
|
||||
mamba_track_interval = server_args.mamba_track_interval
|
||||
|
||||
if len(self.reqs) == 0:
|
||||
self.mamba_track_indices = torch.empty(
|
||||
(0,), dtype=torch.int64, device=self.device
|
||||
)
|
||||
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)
|
||||
set_mamba_track_indices_from_reqs(self)
|
||||
|
||||
|
||||
@@ -357,6 +357,8 @@ class Scheduler(
|
||||
)
|
||||
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
|
||||
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
|
||||
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,
|
||||
tp_group=(
|
||||
self.attn_tp_cpu_group
|
||||
if self.server_args.enable_dp_attention
|
||||
if self.enable_dp_attention
|
||||
else self.tp_cpu_group
|
||||
),
|
||||
tree_cache=self.tree_cache,
|
||||
@@ -911,9 +913,7 @@ class Scheduler(
|
||||
# Use the CPU (gloo) group to broadcast VLM Python objects and avoid CUDA
|
||||
# stream/device coupling (#11910).
|
||||
self.dp_tp_group = (
|
||||
self.attn_tp_group
|
||||
if self.server_args.enable_dp_attention
|
||||
else self.tp_group
|
||||
self.attn_tp_group if self.enable_dp_attention else self.tp_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:
|
||||
# forward_stream is idle (prev forward drained, next not launched),
|
||||
# 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:
|
||||
self.token_to_kv_pool_allocator.flush_opportunistic()
|
||||
except Exception:
|
||||
@@ -3272,7 +3272,7 @@ class Scheduler(
|
||||
self.batch_record_buf[self.batch_record_ct].extend(
|
||||
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
|
||||
# copy_to_cpu); lazy-compaction `_flush` gates src reuse on
|
||||
# it. Only the unified pool's allocator exposes these hooks.
|
||||
@@ -3401,8 +3401,7 @@ class Scheduler(
|
||||
|
||||
def _maybe_report_active_ranks(self) -> None:
|
||||
if not (
|
||||
self.server_args.enable_dp_attention
|
||||
and self.server_args.elastic_ep_backend is not None
|
||||
self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
|
||||
):
|
||||
return
|
||||
# Get the tensors indicating rank activeness
|
||||
@@ -3546,7 +3545,7 @@ class Scheduler(
|
||||
if not self.is_fully_idle():
|
||||
return
|
||||
|
||||
if self.server_args.enable_unified_memory:
|
||||
if self.enable_unified_memory:
|
||||
try:
|
||||
self.token_to_kv_pool_allocator.flush_opportunistic()
|
||||
except Exception:
|
||||
|
||||
@@ -145,6 +145,7 @@ class SchedulerMetricsReporter:
|
||||
self.kv_transfer_latency_ms: float = 0.0
|
||||
|
||||
self.enable_mfu_metrics = False
|
||||
self.decode_log_interval = self.scheduler.server_args.decode_log_interval
|
||||
|
||||
if self.enable_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)
|
||||
|
||||
# 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
|
||||
if (
|
||||
not self.is_stats_logging_rank
|
||||
@@ -716,7 +717,7 @@ class SchedulerMetricsReporter:
|
||||
|
||||
if RECORD_STEP_TIME:
|
||||
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 = (
|
||||
@@ -792,7 +793,7 @@ class SchedulerMetricsReporter:
|
||||
)
|
||||
msg += self._decode_sol_suffix(
|
||||
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_read_bytes = 0.0
|
||||
@@ -1033,10 +1034,7 @@ class SchedulerMetricsReporter:
|
||||
self._device_timer_window_gpu_time / cpu_time * 100, 100
|
||||
)
|
||||
self._device_timer_window_batch_count += 1
|
||||
if (
|
||||
self._device_timer_window_batch_count
|
||||
>= self.scheduler.server_args.decode_log_interval
|
||||
):
|
||||
if self._device_timer_window_batch_count >= self.decode_log_interval:
|
||||
self._device_timer_window_batch_count = 0
|
||||
|
||||
def reset_device_timer_window(self):
|
||||
|
||||
@@ -87,6 +87,7 @@ class _DummyPublisherThread:
|
||||
|
||||
def _fake_server_args(**fields):
|
||||
"""server_args stand-in: carries fields and the override() entry point."""
|
||||
fields.setdefault("decode_log_interval", 40)
|
||||
ns = types.SimpleNamespace(**fields)
|
||||
|
||||
def _override(source, **updates):
|
||||
|
||||
Reference in New Issue
Block a user