[style] Extract init-static values in forward path (#30708)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user