[model_runner] Label forward steps in profile traces with mode and token counts (#23419)
This commit is contained in:
@@ -15,6 +15,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import datetime
|
||||
import gc
|
||||
import inspect
|
||||
@@ -2902,10 +2903,18 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
) -> ModelRunnerOutput:
|
||||
self.forward_pass_id += 1
|
||||
|
||||
with get_global_expert_distribution_recorder().with_forward_pass(
|
||||
self.forward_pass_id,
|
||||
forward_batch,
|
||||
) as recorder_outputs:
|
||||
step_span_ctx = (
|
||||
torch.profiler.record_function(_build_step_span_name(forward_batch))
|
||||
if torch.autograd._profiler_enabled()
|
||||
else contextlib.nullcontext()
|
||||
)
|
||||
with (
|
||||
step_span_ctx,
|
||||
get_global_expert_distribution_recorder().with_forward_pass(
|
||||
self.forward_pass_id,
|
||||
forward_batch,
|
||||
) as recorder_outputs,
|
||||
):
|
||||
output = self._forward_raw(
|
||||
forward_batch,
|
||||
skip_attn_backend_init,
|
||||
@@ -3211,6 +3220,39 @@ def _unwrap_tensor(tensor, tp_rank, device):
|
||||
return tensor.to(device)
|
||||
|
||||
|
||||
def _build_step_span_name(forward_batch: ForwardBatch) -> str:
|
||||
"""Build a profile-trace span name for one forward step.
|
||||
|
||||
Format:
|
||||
step[decode bs=N] — decode-only batch
|
||||
step[prefill bs=N toks=T] — extend-only (prefill) batch
|
||||
step[mixed bs=N ext=T dec=D] — extend+decode mixed batch
|
||||
step[idle] — idle/padding step
|
||||
step[<MODE> bs=N] — other modes (target-verify, etc.)
|
||||
|
||||
Used by ModelRunner.forward to wrap each step in a torch.profile
|
||||
record_function so Chrome traces show labeled step boundaries.
|
||||
"""
|
||||
mode = forward_batch.forward_mode
|
||||
bs = forward_batch.batch_size
|
||||
if mode.is_idle():
|
||||
return "step[idle]"
|
||||
if mode.is_decode():
|
||||
return f"step[decode bs={bs}]"
|
||||
if mode.is_extend():
|
||||
ext_toks = forward_batch.extend_num_tokens or 0
|
||||
ext_seqs = (
|
||||
forward_batch.extend_seq_lens.shape[0]
|
||||
if forward_batch.extend_seq_lens is not None
|
||||
else bs
|
||||
)
|
||||
dec_seqs = bs - ext_seqs
|
||||
if dec_seqs > 0:
|
||||
return f"step[mixed bs={bs} ext={ext_toks} dec={dec_seqs}]"
|
||||
return f"step[prefill bs={bs} toks={ext_toks}]"
|
||||
return f"step[{mode.name} bs={bs}]"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalSerializedTensor:
|
||||
"""torch.Tensor that gets serialized by MultiprocessingSerializer (which only serializes a pointer and not the data).
|
||||
|
||||
Reference in New Issue
Block a user