Add cuda graph status to prefill log (#17836)
This commit is contained in:
@@ -2180,7 +2180,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
forward_batch: ForwardBatch,
|
||||
skip_attn_backend_init: bool = False,
|
||||
pp_proxy_tensors=None,
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
|
||||
) -> Tuple[
|
||||
Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput], bool
|
||||
]:
|
||||
kwargs = {}
|
||||
if self.support_pp:
|
||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||
@@ -2189,20 +2191,28 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if not self.is_generation:
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
if (
|
||||
can_run_graph = (
|
||||
self.piecewise_cuda_graph_runner is not None
|
||||
and self.piecewise_cuda_graph_runner.can_run(forward_batch)
|
||||
):
|
||||
return self.piecewise_cuda_graph_runner.replay(forward_batch, **kwargs)
|
||||
)
|
||||
|
||||
if can_run_graph:
|
||||
return (
|
||||
self.piecewise_cuda_graph_runner.replay(forward_batch, **kwargs),
|
||||
can_run_graph,
|
||||
)
|
||||
|
||||
if not skip_attn_backend_init:
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
|
||||
return self.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
return (
|
||||
self.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
),
|
||||
can_run_graph,
|
||||
)
|
||||
|
||||
def forward_idle(
|
||||
@@ -2358,7 +2368,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
forward_count=split_forward_count,
|
||||
)
|
||||
elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
|
||||
ret = self.forward_extend(
|
||||
ret, can_run_graph = self.forward_extend(
|
||||
forward_batch,
|
||||
skip_attn_backend_init=skip_attn_backend_init,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
|
||||
Reference in New Issue
Block a user