[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
|
# send out proxy tensors to the next stage
|
||||||
if self.cur_batch:
|
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(
|
self.pp_group.send_tensor_dict(
|
||||||
result.pp_hidden_states_proxy_tensors.tensors,
|
result.pp_hidden_states_proxy_tensors.tensors,
|
||||||
all_gather_group=self.attn_tp_group,
|
all_gather_group=self.attn_tp_group,
|
||||||
|
|||||||
@@ -327,6 +327,11 @@ class ModelRunner:
|
|||||||
"pp_proxy_tensors" in inspect.signature(self.model.forward).parameters
|
"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
|
# For weight updates
|
||||||
self._model_update_group = {}
|
self._model_update_group = {}
|
||||||
self._weights_send_group = {}
|
self._weights_send_group = {}
|
||||||
@@ -2056,7 +2061,7 @@ class ModelRunner:
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
skip_attn_backend_init: bool = False,
|
skip_attn_backend_init: bool = False,
|
||||||
pp_proxy_tensors=None,
|
pp_proxy_tensors=None,
|
||||||
) -> LogitsProcessorOutput:
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
if not skip_attn_backend_init:
|
if not skip_attn_backend_init:
|
||||||
if self.server_args.enable_pdmux:
|
if self.server_args.enable_pdmux:
|
||||||
self.decode_attn_backend.init_forward_metadata(forward_batch)
|
self.decode_attn_backend.init_forward_metadata(forward_batch)
|
||||||
@@ -2079,7 +2084,7 @@ class ModelRunner:
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
skip_attn_backend_init: bool = False,
|
skip_attn_backend_init: bool = False,
|
||||||
pp_proxy_tensors=None,
|
pp_proxy_tensors=None,
|
||||||
) -> LogitsProcessorOutput:
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
kwargs = {}
|
kwargs = {}
|
||||||
if self.support_pp:
|
if self.support_pp:
|
||||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||||
@@ -2106,7 +2111,7 @@ class ModelRunner:
|
|||||||
|
|
||||||
def forward_idle(
|
def forward_idle(
|
||||||
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
||||||
) -> LogitsProcessorOutput:
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
kwargs = {}
|
kwargs = {}
|
||||||
if self.support_pp:
|
if self.support_pp:
|
||||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||||
|
|||||||
Reference in New Issue
Block a user