Add cuda graph status to prefill log (#17836)

This commit is contained in:
Ke Bao
2026-01-30 16:56:53 +08:00
committed by GitHub
parent c8dc543dc5
commit 77a27e728c
4 changed files with 62 additions and 25 deletions
@@ -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,