[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):
|
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)
|
||||||
|
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user