diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 7553f4702..13a1bdf7c 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 5669ab2dc..d464e74af 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index ff48bce8c..f82face7c 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -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): diff --git a/test/registered/unit/observability/test_forward_pass_metrics.py b/test/registered/unit/observability/test_forward_pass_metrics.py index ae24e80b6..56cd4f70b 100644 --- a/test/registered/unit/observability/test_forward_pass_metrics.py +++ b/test/registered/unit/observability/test_forward_pass_metrics.py @@ -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):