Custom All Reduce for Piecewise Cuda Graph (#15356)
Signed-off-by: Oasis-Git <ayw.sirius19@gmail.com>
This commit is contained in:
@@ -13,7 +13,6 @@ import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.distributed.parallel_state import get_tp_group
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.multimodal import gpu_tensor_hash
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
@@ -1043,69 +1042,52 @@ def general_mm_embed_routine(
|
||||
Returns:
|
||||
Hidden states from language model forward pass
|
||||
"""
|
||||
# Lazy import to allow some monkey patch of piecewise_cuda_graph_runner
|
||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||
use_original_ca_comm,
|
||||
)
|
||||
|
||||
tp_group = get_tp_group()
|
||||
|
||||
with use_original_ca_comm(tp_group):
|
||||
# We disable custom allreduce in piecewise cuda graph.
|
||||
# However, because we only capture the language model part, the multimodal can still use custom allreduce.
|
||||
assert hasattr(language_model, "get_input_embeddings")
|
||||
embed_tokens = language_model.get_input_embeddings()
|
||||
assert hasattr(language_model, "get_input_embeddings")
|
||||
embed_tokens = language_model.get_input_embeddings()
|
||||
if not hasattr(language_model, "pp_group") or language_model.pp_group.is_first_rank:
|
||||
if (
|
||||
not hasattr(language_model, "pp_group")
|
||||
or language_model.pp_group.is_first_rank
|
||||
not forward_batch.forward_mode.is_decode()
|
||||
and not forward_batch.forward_mode.is_target_verify()
|
||||
and forward_batch.contains_mm_inputs()
|
||||
):
|
||||
if (
|
||||
not forward_batch.forward_mode.is_decode()
|
||||
and not forward_batch.forward_mode.is_target_verify()
|
||||
and forward_batch.contains_mm_inputs()
|
||||
):
|
||||
mm_inputs_list = [
|
||||
mm_input
|
||||
for mm_input in forward_batch.mm_inputs
|
||||
if mm_input is not None
|
||||
]
|
||||
extend_prefix_lens = [
|
||||
prefix_len
|
||||
for i, prefix_len in enumerate(forward_batch.extend_prefix_lens_cpu)
|
||||
if forward_batch.mm_inputs[i] is not None
|
||||
]
|
||||
extend_seq_lens = [
|
||||
seq_len
|
||||
for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
|
||||
if forward_batch.mm_inputs[i] is not None
|
||||
]
|
||||
input_embeds, other_info = embed_mm_inputs(
|
||||
mm_inputs_list=mm_inputs_list,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
input_ids=input_ids,
|
||||
multimodal_model=multimodal_model,
|
||||
input_embedding=embed_tokens,
|
||||
data_embedding_func_mapping=data_embedding_funcs,
|
||||
placeholder_tokens=placeholder_tokens,
|
||||
use_deepstack=use_deepstack,
|
||||
)
|
||||
# add for qwen3_vl deepstack
|
||||
if use_deepstack:
|
||||
kwargs["input_deepstack_embeds"] = other_info[
|
||||
"input_deepstack_embeds"
|
||||
]
|
||||
# once used, mm_inputs is useless, considering chunked-prefill is disabled for multimodal models
|
||||
# just being defensive here
|
||||
forward_batch.mm_inputs = None
|
||||
else:
|
||||
input_embeds = embed_tokens(input_ids)
|
||||
# Copy to pre-allocated buffer if available (for CUDA graph address stability)
|
||||
if forward_batch.input_embeds is not None:
|
||||
forward_batch.input_embeds.copy_(input_embeds)
|
||||
input_embeds = forward_batch.input_embeds
|
||||
mm_inputs_list = [
|
||||
mm_input for mm_input in forward_batch.mm_inputs if mm_input is not None
|
||||
]
|
||||
extend_prefix_lens = [
|
||||
prefix_len
|
||||
for i, prefix_len in enumerate(forward_batch.extend_prefix_lens_cpu)
|
||||
if forward_batch.mm_inputs[i] is not None
|
||||
]
|
||||
extend_seq_lens = [
|
||||
seq_len
|
||||
for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
|
||||
if forward_batch.mm_inputs[i] is not None
|
||||
]
|
||||
input_embeds, other_info = embed_mm_inputs(
|
||||
mm_inputs_list=mm_inputs_list,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
input_ids=input_ids,
|
||||
multimodal_model=multimodal_model,
|
||||
input_embedding=embed_tokens,
|
||||
data_embedding_func_mapping=data_embedding_funcs,
|
||||
placeholder_tokens=placeholder_tokens,
|
||||
use_deepstack=use_deepstack,
|
||||
)
|
||||
# add for qwen3_vl deepstack
|
||||
if use_deepstack:
|
||||
kwargs["input_deepstack_embeds"] = other_info["input_deepstack_embeds"]
|
||||
# once used, mm_inputs is useless, considering chunked-prefill is disabled for multimodal models
|
||||
# just being defensive here
|
||||
forward_batch.mm_inputs = None
|
||||
else:
|
||||
input_embeds = None
|
||||
input_embeds = embed_tokens(input_ids)
|
||||
# Copy to pre-allocated buffer if available (for CUDA graph address stability)
|
||||
if forward_batch.input_embeds is not None:
|
||||
forward_batch.input_embeds.copy_(input_embeds)
|
||||
input_embeds = forward_batch.input_embeds
|
||||
else:
|
||||
input_embeds = None
|
||||
|
||||
hidden_states = language_model(
|
||||
input_ids=None,
|
||||
|
||||
Reference in New Issue
Block a user