Custom All Reduce for Piecewise Cuda Graph (#15356)

Signed-off-by: Oasis-Git <ayw.sirius19@gmail.com>
This commit is contained in:
Yuwei An
2025-12-25 23:55:15 +08:00
committed by GitHub
parent b6702d72cf
commit 5c243ba588
5 changed files with 146 additions and 145 deletions
+43 -61
View File
@@ -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,