[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.data_type = model_runner.kv_cache_dtype
|
||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
self.req_to_token_pool = model_runner.req_to_token_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:
|
def _select_backend(self, forward_mode: ForwardMode) -> AttentionBackend:
|
||||||
"""
|
"""
|
||||||
@@ -45,7 +51,7 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
return (
|
return (
|
||||||
self.decode_backend
|
self.decode_backend
|
||||||
if self.model_runner.server_args.speculative_attention_mode == "decode"
|
if self.spec_attn_is_decode
|
||||||
else self.prefill_backend
|
else self.prefill_backend
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -71,7 +77,7 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||||
if (
|
if (
|
||||||
self.model_runner.server_args.speculative_algorithm is not None
|
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
|
# When speculative decoding is enabled, we need to initialize the backend
|
||||||
# that will be used for target_verify.
|
# that will be used for target_verify.
|
||||||
@@ -144,7 +150,7 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
return backend.get_indexer_metadata(layer_id, forward_batch)
|
return backend.get_indexer_metadata(layer_id, forward_batch)
|
||||||
|
|
||||||
def update_mamba_state_after_mtp_verify(self, *args, **kwargs):
|
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
|
backend = self.decode_backend
|
||||||
else:
|
else:
|
||||||
backend = self.prefill_backend
|
backend = self.prefill_backend
|
||||||
|
|||||||
@@ -360,6 +360,7 @@ class LogitsProcessor(nn.Module):
|
|||||||
|
|
||||||
self.return_full_logits = return_full_logits
|
self.return_full_logits = return_full_logits
|
||||||
self.enable_mis = get_server_args().enable_mis
|
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(
|
self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
|
||||||
max_tokens=triton_symm_mem_ag.recommended_max_tokens(
|
max_tokens=triton_symm_mem_ag.recommended_max_tokens(
|
||||||
@@ -969,7 +970,7 @@ class LogitsProcessor(nn.Module):
|
|||||||
None, # bias
|
None, # bias
|
||||||
True, # is_vnni
|
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
|
# Due to tie-weight, we may not be able to change lm_head's weight dtype
|
||||||
logits = torch.matmul(
|
logits = torch.matmul(
|
||||||
hidden_states.bfloat16(), lm_head.weight.T.bfloat16()
|
hidden_states.bfloat16(), lm_head.weight.T.bfloat16()
|
||||||
|
|||||||
@@ -553,6 +553,9 @@ class CPUGraphRunner:
|
|||||||
# Parse args
|
# Parse args
|
||||||
self.model_runner = model_runner
|
self.model_runner = model_runner
|
||||||
self.device = model_runner.device
|
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)
|
# bs -> compiled fn (text-only / skip_cross_attention=True)
|
||||||
self.graphs = {}
|
self.graphs = {}
|
||||||
# bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only)
|
# bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only)
|
||||||
@@ -581,7 +584,7 @@ class CPUGraphRunner:
|
|||||||
self.num_tokens_per_bs = 1
|
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 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
|
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
@@ -870,7 +873,7 @@ class CPUGraphRunner:
|
|||||||
)
|
)
|
||||||
capture_hidden_mode_required_for_returning_hidden_states = (
|
capture_hidden_mode_required_for_returning_hidden_states = (
|
||||||
CaptureHiddenMode.FULL
|
CaptureHiddenMode.FULL
|
||||||
if self.model_runner.server_args.enable_return_hidden_states
|
if self.enable_return_hidden_states
|
||||||
else CaptureHiddenMode.NULL
|
else CaptureHiddenMode.NULL
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1078,14 +1078,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# batch_size * [3 * seq_len]
|
# batch_size * [3 * seq_len]
|
||||||
batch_size = self.seq_lens_cpu.shape[0]
|
batch_size = self.seq_lens_cpu.shape[0]
|
||||||
mrope_positions_list = [[]] * batch_size
|
mrope_positions_list = [[]] * batch_size
|
||||||
|
rl_on_policy_target = get_server_args().rl_on_policy_target
|
||||||
for batch_idx in range(batch_size):
|
for batch_idx in range(batch_size):
|
||||||
mm_input = batch.multimodal_inputs[batch_idx]
|
mm_input = batch.multimodal_inputs[batch_idx]
|
||||||
if self.forward_mode.is_decode():
|
if self.forward_mode.is_decode():
|
||||||
# 3 * N
|
# 3 * N
|
||||||
if (
|
if mm_input is None or rl_on_policy_target is not None:
|
||||||
mm_input is None
|
|
||||||
or get_server_args().rl_on_policy_target is not None
|
|
||||||
):
|
|
||||||
mrope_positions_list[batch_idx] = torch.full(
|
mrope_positions_list[batch_idx] = torch.full(
|
||||||
(3, 1),
|
(3, 1),
|
||||||
self.seq_lens_cpu[batch_idx] - 1,
|
self.seq_lens_cpu[batch_idx] - 1,
|
||||||
@@ -1101,10 +1099,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
batch.extend_lens[batch_idx],
|
batch.extend_lens[batch_idx],
|
||||||
batch.prefix_lens[batch_idx],
|
batch.prefix_lens[batch_idx],
|
||||||
)
|
)
|
||||||
if (
|
if mm_input is None or rl_on_policy_target is not None:
|
||||||
mm_input is None
|
|
||||||
or get_server_args().rl_on_policy_target is not None
|
|
||||||
):
|
|
||||||
# text only
|
# text only
|
||||||
mrope_positions = torch.tensor(
|
mrope_positions = torch.tensor(
|
||||||
[
|
[
|
||||||
|
|||||||
@@ -3050,7 +3050,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.msprobe_debugger.stop()
|
self.msprobe_debugger.stop()
|
||||||
self.msprobe_debugger.step()
|
self.msprobe_debugger.step()
|
||||||
|
|
||||||
if self.server_args.elastic_ep_backend is not None:
|
if self.enable_elastic_ep:
|
||||||
self.maybe_recover_ep_ranks()
|
self.maybe_recover_ep_ranks()
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|||||||
@@ -199,6 +199,10 @@ class BaseRunner(ABC):
|
|||||||
self.tp_size = model_runner.server_args.tp_size
|
self.tp_size = model_runner.server_args.tp_size
|
||||||
self.dp_size = model_runner.server_args.dp_size
|
self.dp_size = model_runner.server_args.dp_size
|
||||||
self.pp_size = model_runner.server_args.pp_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_size = get_parallel().attn_tp_size
|
||||||
self.attn_tp_rank = get_parallel().attn_tp_rank
|
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||||
|
|||||||
@@ -204,7 +204,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.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_size = get_parallel().attn_tp_size
|
||||||
self.attn_tp_rank = get_parallel().attn_tp_rank
|
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||||
@@ -259,7 +258,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
KTMoEWrapper.set_capture_batch_sizes(self.capture_bs)
|
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 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
|
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
# Attention backend
|
# Attention backend
|
||||||
@@ -882,7 +881,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
)
|
)
|
||||||
capture_hidden_mode_required_for_returning_hidden_states = (
|
capture_hidden_mode_required_for_returning_hidden_states = (
|
||||||
CaptureHiddenMode.FULL
|
CaptureHiddenMode.FULL
|
||||||
if self.model_runner.server_args.enable_return_hidden_states
|
if self.enable_return_hidden_states
|
||||||
else CaptureHiddenMode.NULL
|
else CaptureHiddenMode.NULL
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -216,7 +216,7 @@ class EagerRunner(BaseRunner):
|
|||||||
runs under. PDmux selects a per-stream backend and publishes it via an
|
runs under. PDmux selects a per-stream backend and publishes it via an
|
||||||
active ForwardContext; non-pdmux uses attn_backend + the ambient ctx."""
|
active ForwardContext; non-pdmux uses attn_backend + the ambient ctx."""
|
||||||
model_runner = self.model_runner
|
model_runner = self.model_runner
|
||||||
if model_runner.server_args.enable_pdmux:
|
if self.enable_pdmux:
|
||||||
return model_runner.decode_attn_backend, forward_context(
|
return model_runner.decode_attn_backend, forward_context(
|
||||||
ForwardContext(attn_backend=model_runner.decode_attn_backend)
|
ForwardContext(attn_backend=model_runner.decode_attn_backend)
|
||||||
)
|
)
|
||||||
@@ -228,7 +228,7 @@ class EagerRunner(BaseRunner):
|
|||||||
pp_proxy_tensors=None,
|
pp_proxy_tensors=None,
|
||||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
model_runner = self.model_runner
|
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()
|
attn_backend, pdmux_ctx = self._resolve_decode_pdmux()
|
||||||
if not enable_pdmux:
|
if not enable_pdmux:
|
||||||
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
@@ -263,7 +263,7 @@ class EagerRunner(BaseRunner):
|
|||||||
model_runner = self.model_runner
|
model_runner = self.model_runner
|
||||||
kwargs = model_runner._extend_forward_kwargs(forward_batch, pp_proxy_tensors)
|
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)
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
|
|
||||||
if forward_batch.needs_forward_metadata_init():
|
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
|
# 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.
|
# drop stale forward_metadata to avoid an SWA use-after-free on req_pool.
|
||||||
if forward_batch.batch_size > 0:
|
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)
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user