[Feature] PD-Multiplexing Context and Scheduler, lazy import spatial. (#12275)

This commit is contained in:
ykcombat
2025-11-01 00:40:01 +08:00
committed by GitHub
parent 70562969b9
commit 41efcaeb45
9 changed files with 458 additions and 24 deletions

View File

@@ -96,6 +96,7 @@ class ForwardMode(IntEnum):
else False
)
or self == ForwardMode.TARGET_VERIFY
or self == ForwardMode.SPLIT_PREFILL
)
def is_decode(self):

View File

@@ -1765,6 +1765,7 @@ class ModelRunner:
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.start_layer,
end_layer=self.end_layer,
enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(
self.server_args.speculative_algorithm is not None
),
@@ -1833,12 +1834,18 @@ class ModelRunner:
def init_attention_backend(self):
"""Init attention kernel backend."""
if self.server_args.enable_two_batch_overlap and not self.is_draft_worker:
if self.server_args.enable_pdmux:
self.attn_backend = self._get_attention_backend(init_new_workspace=True)
self.decode_attn_backend_group = []
for _ in range(self.server_args.sm_group_num):
self.decode_attn_backend_group.append(self._get_attention_backend())
self.decode_attn_backend = self.decode_attn_backend_group[0]
elif self.server_args.enable_two_batch_overlap and not self.is_draft_worker:
self.attn_backend = TboAttnBackend.init_new(self._get_attention_backend)
else:
self.attn_backend = self._get_attention_backend()
def _get_attention_backend(self):
def _get_attention_backend(self, init_new_workspace: bool = False):
"""Init attention kernel backend."""
self.prefill_attention_backend_str, self.decode_attention_backend_str = (
self.server_args.get_attention_backends()
@@ -1852,10 +1859,12 @@ class ModelRunner:
attn_backend = HybridAttnBackend(
self,
decode_backend=self._get_attention_backend_from_str(
self.decode_attention_backend_str
self.decode_attention_backend_str,
init_new_workspace=init_new_workspace,
),
prefill_backend=self._get_attention_backend_from_str(
self.prefill_attention_backend_str
self.prefill_attention_backend_str,
init_new_workspace=init_new_workspace,
),
)
logger.info(
@@ -1869,7 +1878,8 @@ class ModelRunner:
)
else:
attn_backend = self._get_attention_backend_from_str(
self.server_args.attention_backend
self.server_args.attention_backend,
init_new_workspace=init_new_workspace,
)
(
@@ -1878,9 +1888,12 @@ class ModelRunner:
) = (self.prefill_attention_backend_str, self.decode_attention_backend_str)
return attn_backend
def _get_attention_backend_from_str(self, backend_str: str):
def _get_attention_backend_from_str(
self, backend_str: str, init_new_workspace: bool = False
):
if backend_str not in ATTENTION_BACKENDS:
raise ValueError(f"Invalid attention backend: {backend_str}")
self.init_new_workspace = init_new_workspace
full_attention_backend = ATTENTION_BACKENDS[backend_str](self)
return attn_backend_wrapper(self, full_attention_backend)
@@ -1978,6 +1991,9 @@ class ModelRunner:
device_mesh = torch.distributed.init_device_mesh(self.device, (self.tp_size,))
tensor_parallel(self.model, device_mesh)
def update_decode_attn_backend(self, stream_idx: int):
self.decode_attn_backend = self.decode_attn_backend_group[stream_idx]
def forward_decode(
self,
forward_batch: ForwardBatch,
@@ -1985,7 +2001,11 @@ class ModelRunner:
pp_proxy_tensors=None,
) -> LogitsProcessorOutput:
if not skip_attn_backend_init:
self.attn_backend.init_forward_metadata(forward_batch)
if self.server_args.enable_pdmux:
self.decode_attn_backend.init_forward_metadata(forward_batch)
forward_batch.attn_backend = self.decode_attn_backend
else:
self.attn_backend.init_forward_metadata(forward_batch)
# FIXME: add pp_proxy_tensors arg to all models
kwargs = {}
if self.support_pp:
@@ -2123,18 +2143,18 @@ class ModelRunner:
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_extend():
ret = self.forward_extend(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_split_prefill():
ret = self.forward_split_prefill(
forward_batch,
reinit_attn_backend=reinit_attn_backend,
forward_count=split_forward_count,
)
elif forward_batch.forward_mode.is_extend():
ret = self.forward_extend(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_idle():
ret = self.forward_idle(forward_batch, pp_proxy_tensors=pp_proxy_tensors)
else: