[PP] put pp assert in model runner (#12934)
Signed-off-by: Xuchun Shang <xuchun.shang@gmail.com>
This commit is contained in:
@@ -122,8 +122,6 @@ class SchedulerPPMixin:
|
||||
|
||||
# send out proxy tensors to the next stage
|
||||
if self.cur_batch:
|
||||
# FIXME(lsyin): remove this assert
|
||||
assert result.pp_hidden_states_proxy_tensors.tensors is not None
|
||||
self.pp_group.send_tensor_dict(
|
||||
result.pp_hidden_states_proxy_tensors.tensors,
|
||||
all_gather_group=self.attn_tp_group,
|
||||
|
||||
@@ -327,6 +327,11 @@ class ModelRunner:
|
||||
"pp_proxy_tensors" in inspect.signature(self.model.forward).parameters
|
||||
)
|
||||
|
||||
if self.pp_size > 1:
|
||||
assert (
|
||||
self.support_pp
|
||||
), "Pipeline Parallel is not compatible with this model."
|
||||
|
||||
# For weight updates
|
||||
self._model_update_group = {}
|
||||
self._weights_send_group = {}
|
||||
@@ -2056,7 +2061,7 @@ class ModelRunner:
|
||||
forward_batch: ForwardBatch,
|
||||
skip_attn_backend_init: bool = False,
|
||||
pp_proxy_tensors=None,
|
||||
) -> LogitsProcessorOutput:
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||
if not skip_attn_backend_init:
|
||||
if self.server_args.enable_pdmux:
|
||||
self.decode_attn_backend.init_forward_metadata(forward_batch)
|
||||
@@ -2079,7 +2084,7 @@ class ModelRunner:
|
||||
forward_batch: ForwardBatch,
|
||||
skip_attn_backend_init: bool = False,
|
||||
pp_proxy_tensors=None,
|
||||
) -> LogitsProcessorOutput:
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||
kwargs = {}
|
||||
if self.support_pp:
|
||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||
@@ -2106,7 +2111,7 @@ class ModelRunner:
|
||||
|
||||
def forward_idle(
|
||||
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
||||
) -> LogitsProcessorOutput:
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||
kwargs = {}
|
||||
if self.support_pp:
|
||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||
|
||||
Reference in New Issue
Block a user