profile: add vlm prefill profiler ranges (#30871)

This commit is contained in:
Mick
2026-07-12 14:07:10 +08:00
committed by GitHub
parent bce3fc987d
commit f1c247edf9
+37 -32
View File
@@ -1074,32 +1074,36 @@ def general_mm_embed_routine(
if forward_batch.mm_inputs[i] is not None if forward_batch.mm_inputs[i] is not None
] ]
server_args = get_server_args() server_args = get_server_args()
if server_args and server_args.enable_adaptive_dispatch_to_encoder: # Makes VLM profiles directly attributable: this range includes
# Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches # encoder/ViT execution and multimodal feature placement, while
input_embeds, other_info = _embed_mm_inputs_with_split( # the language model range below excludes both.
mm_inputs_list=mm_inputs_list, with torch.profiler.record_function("sglang.vlm.mm_embedding"):
extend_prefix_lens=extend_prefix_lens, if server_args and server_args.enable_adaptive_dispatch_to_encoder:
extend_seq_lens=extend_seq_lens, # Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
input_ids=input_ids, input_embeds, other_info = _embed_mm_inputs_with_split(
forward_batch=forward_batch, mm_inputs_list=mm_inputs_list,
input_embedding=embed_tokens, extend_prefix_lens=extend_prefix_lens,
multimodal_model=multimodal_model, extend_seq_lens=extend_seq_lens,
data_embedding_func_mapping=data_embedding_funcs, input_ids=input_ids,
placeholder_tokens=placeholder_tokens, forward_batch=forward_batch,
use_deepstack=use_deepstack, input_embedding=embed_tokens,
) multimodal_model=multimodal_model,
else: data_embedding_func_mapping=data_embedding_funcs,
input_embeds, other_info = embed_mm_inputs( placeholder_tokens=placeholder_tokens,
mm_inputs_list=mm_inputs_list, use_deepstack=use_deepstack,
extend_prefix_lens=extend_prefix_lens, )
extend_seq_lens=extend_seq_lens, else:
input_ids=input_ids, input_embeds, other_info = embed_mm_inputs(
input_embedding=embed_tokens, mm_inputs_list=mm_inputs_list,
multimodal_model=multimodal_model, extend_prefix_lens=extend_prefix_lens,
data_embedding_func_mapping=data_embedding_funcs, extend_seq_lens=extend_seq_lens,
placeholder_tokens=placeholder_tokens, input_ids=input_ids,
use_deepstack=use_deepstack, input_embedding=embed_tokens,
) multimodal_model=multimodal_model,
data_embedding_func_mapping=data_embedding_funcs,
placeholder_tokens=placeholder_tokens,
use_deepstack=use_deepstack,
)
# add for qwen3_vl deepstack # add for qwen3_vl deepstack
if use_deepstack: if use_deepstack:
@@ -1143,12 +1147,13 @@ def general_mm_embed_routine(
else: else:
input_embeds = None input_embeds = None
hidden_states = language_model( with torch.profiler.record_function("sglang.vlm.language_model_prefill"):
input_ids=None, hidden_states = language_model(
forward_batch=forward_batch, input_ids=None,
input_embeds=input_embeds, forward_batch=forward_batch,
**kwargs, input_embeds=input_embeds,
) **kwargs,
)
return hidden_states return hidden_states