diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index af5bd4de2..40f775d8f 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -24,6 +24,12 @@ class HybridAttnBackend(AttentionBackend): self.data_type = model_runner.kv_cache_dtype self.token_to_kv_pool = model_runner.token_to_kv_pool self.req_to_token_pool = model_runner.req_to_token_pool + self.spec_attn_is_decode = ( + model_runner.server_args.speculative_attention_mode == "decode" + ) + self.spec_attn_is_prefill = ( + model_runner.server_args.speculative_attention_mode == "prefill" + ) def _select_backend(self, forward_mode: ForwardMode) -> AttentionBackend: """ @@ -45,7 +51,7 @@ class HybridAttnBackend(AttentionBackend): elif forward_mode.is_target_verify(): return ( self.decode_backend - if self.model_runner.server_args.speculative_attention_mode == "decode" + if self.spec_attn_is_decode else self.prefill_backend ) else: @@ -71,7 +77,7 @@ class HybridAttnBackend(AttentionBackend): self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens) if ( self.model_runner.server_args.speculative_algorithm is not None - and self.model_runner.server_args.speculative_attention_mode == "prefill" + and self.spec_attn_is_prefill ): # When speculative decoding is enabled, we need to initialize the backend # that will be used for target_verify. @@ -144,7 +150,7 @@ class HybridAttnBackend(AttentionBackend): return backend.get_indexer_metadata(layer_id, forward_batch) def update_mamba_state_after_mtp_verify(self, *args, **kwargs): - if self.model_runner.server_args.speculative_attention_mode == "decode": + if self.spec_attn_is_decode: backend = self.decode_backend else: backend = self.prefill_backend diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index a413b1836..aad02844a 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -360,6 +360,7 @@ class LogitsProcessor(nn.Module): self.return_full_logits = return_full_logits self.enable_mis = get_server_args().enable_mis + self.rl_on_policy_target = get_server_args().rl_on_policy_target self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer( max_tokens=triton_symm_mem_ag.recommended_max_tokens( @@ -969,7 +970,7 @@ class LogitsProcessor(nn.Module): None, # bias True, # is_vnni ) - elif get_server_args().rl_on_policy_target is not None: + elif self.rl_on_policy_target is not None: # Due to tie-weight, we may not be able to change lm_head's weight dtype logits = torch.matmul( hidden_states.bfloat16(), lm_head.weight.T.bfloat16() diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 12405a059..b058d8fb0 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -553,6 +553,9 @@ class CPUGraphRunner: # Parse args self.model_runner = model_runner self.device = model_runner.device + self.enable_return_hidden_states = ( + model_runner.server_args.enable_return_hidden_states + ) # bs -> compiled fn (text-only / skip_cross_attention=True) self.graphs = {} # bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only) @@ -581,7 +584,7 @@ class CPUGraphRunner: self.num_tokens_per_bs = 1 # If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup - if model_runner.server_args.enable_return_hidden_states: + if self.enable_return_hidden_states: self.capture_hidden_mode = CaptureHiddenMode.FULL assert ( @@ -870,7 +873,7 @@ class CPUGraphRunner: ) capture_hidden_mode_required_for_returning_hidden_states = ( CaptureHiddenMode.FULL - if self.model_runner.server_args.enable_return_hidden_states + if self.enable_return_hidden_states else CaptureHiddenMode.NULL ) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index c013d3d56..e1fb9b28e 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -1078,14 +1078,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # batch_size * [3 * seq_len] batch_size = self.seq_lens_cpu.shape[0] mrope_positions_list = [[]] * batch_size + rl_on_policy_target = get_server_args().rl_on_policy_target for batch_idx in range(batch_size): mm_input = batch.multimodal_inputs[batch_idx] if self.forward_mode.is_decode(): # 3 * N - if ( - mm_input is None - or get_server_args().rl_on_policy_target is not None - ): + if mm_input is None or rl_on_policy_target is not None: mrope_positions_list[batch_idx] = torch.full( (3, 1), self.seq_lens_cpu[batch_idx] - 1, @@ -1101,10 +1099,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): batch.extend_lens[batch_idx], batch.prefix_lens[batch_idx], ) - if ( - mm_input is None - or get_server_args().rl_on_policy_target is not None - ): + if mm_input is None or rl_on_policy_target is not None: # text only mrope_positions = torch.tensor( [ diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 4cd22513e..ca641d1fe 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -3050,7 +3050,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.msprobe_debugger.stop() self.msprobe_debugger.step() - if self.server_args.elastic_ep_backend is not None: + if self.enable_elastic_ep: self.maybe_recover_ep_ranks() return output diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index edc9e8a61..56420661a 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -199,6 +199,10 @@ class BaseRunner(ABC): self.tp_size = model_runner.server_args.tp_size self.dp_size = model_runner.server_args.dp_size self.pp_size = model_runner.server_args.pp_size + self.enable_pdmux = model_runner.server_args.enable_pdmux + self.enable_return_hidden_states = ( + model_runner.server_args.enable_return_hidden_states + ) self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_rank = get_parallel().attn_tp_rank self.tbo_plugin = TboCudaGraphRunnerPlugin() diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index ea7897e79..83a1b4ac5 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -204,7 +204,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) - self.enable_pdmux = model_runner.server_args.enable_pdmux self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_rank = get_parallel().attn_tp_rank @@ -259,7 +258,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): KTMoEWrapper.set_capture_batch_sizes(self.capture_bs) # If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup - if model_runner.server_args.enable_return_hidden_states: + if self.enable_return_hidden_states: self.capture_hidden_mode = CaptureHiddenMode.FULL # Attention backend @@ -882,7 +881,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): ) capture_hidden_mode_required_for_returning_hidden_states = ( CaptureHiddenMode.FULL - if self.model_runner.server_args.enable_return_hidden_states + if self.enable_return_hidden_states else CaptureHiddenMode.NULL ) diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 52a76fdb3..cf583bab2 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -216,7 +216,7 @@ class EagerRunner(BaseRunner): runs under. PDmux selects a per-stream backend and publishes it via an active ForwardContext; non-pdmux uses attn_backend + the ambient ctx.""" model_runner = self.model_runner - if model_runner.server_args.enable_pdmux: + if self.enable_pdmux: return model_runner.decode_attn_backend, forward_context( ForwardContext(attn_backend=model_runner.decode_attn_backend) ) @@ -228,7 +228,7 @@ class EagerRunner(BaseRunner): pp_proxy_tensors=None, ) -> Union[LogitsProcessorOutput, PPProxyTensors]: model_runner = self.model_runner - enable_pdmux = model_runner.server_args.enable_pdmux + enable_pdmux = self.enable_pdmux attn_backend, pdmux_ctx = self._resolve_decode_pdmux() if not enable_pdmux: forward_batch = self.load_batch(forward_batch, pp_proxy_tensors) @@ -263,7 +263,7 @@ class EagerRunner(BaseRunner): model_runner = self.model_runner kwargs = model_runner._extend_forward_kwargs(forward_batch, pp_proxy_tensors) - if not model_runner.server_args.enable_pdmux: + if not self.enable_pdmux: forward_batch = self.load_batch(forward_batch, pp_proxy_tensors) if forward_batch.needs_forward_metadata_init(): @@ -393,7 +393,7 @@ class EagerRunner(BaseRunner): # Padded idle (DP-attn MLP sync) needs metadata reinit; unpadded must # drop stale forward_metadata to avoid an SWA use-after-free on req_pool. if forward_batch.batch_size > 0: - if not model_runner.server_args.enable_pdmux: + if not self.enable_pdmux: forward_batch = self.load_batch(forward_batch, pp_proxy_tensors) model_runner.attn_backend.init_forward_metadata(forward_batch) else: