DP Enhancement (#8280)
This commit is contained in:
@@ -1464,9 +1464,13 @@ class ModelRunner:
|
||||
tensor_parallel(self.model, device_mesh)
|
||||
|
||||
def forward_decode(
|
||||
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
skip_attn_backend_init: bool = False,
|
||||
pp_proxy_tensors=None,
|
||||
) -> LogitsProcessorOutput:
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
if not skip_attn_backend_init:
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
# FIXME: add pp_proxy_tensors arg to all models
|
||||
kwargs = {}
|
||||
if self.support_pp:
|
||||
@@ -1578,8 +1582,18 @@ class ModelRunner:
|
||||
skip_attn_backend_init=skip_attn_backend_init,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
)
|
||||
elif forward_batch.forward_mode.is_decode():
|
||||
ret = self.forward_decode(forward_batch, pp_proxy_tensors=pp_proxy_tensors)
|
||||
return ret, can_run_cuda_graph
|
||||
|
||||
# For MLP sync
|
||||
if forward_batch.global_num_tokens_cpu is not None:
|
||||
forward_batch.prepare_mlp_sync_batch(self)
|
||||
|
||||
if forward_batch.forward_mode.is_decode():
|
||||
ret = self.forward_decode(
|
||||
forward_batch,
|
||||
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,
|
||||
@@ -1597,6 +1611,9 @@ class ModelRunner:
|
||||
else:
|
||||
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode}")
|
||||
|
||||
if forward_batch.global_num_tokens_cpu is not None:
|
||||
forward_batch.post_forward_mlp_sync_batch(ret)
|
||||
|
||||
return ret, can_run_cuda_graph
|
||||
|
||||
def _preprocess_logits(
|
||||
|
||||
Reference in New Issue
Block a user