[style] Extract init-static values in forward path (#30708)

This commit is contained in:
Liangsheng Yin
2026-07-09 19:35:00 -07:00
committed by GitHub
parent b5e75b9423
commit dda61b476e
8 changed files with 30 additions and 22 deletions
@@ -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
+2 -1
View File
@@ -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()
@@ -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
)
@@ -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(
[
@@ -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
@@ -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()
@@ -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
)
@@ -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: