[Feature] Enable CUDA graph for PD-Multiplexing. (#11595)

This commit is contained in:
ykcombat
2025-11-14 03:39:40 +08:00
committed by GitHub
parent bfe638f7e8
commit dd192a55f4
3 changed files with 103 additions and 52 deletions

View File

@@ -34,7 +34,11 @@ from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
set_graph_pool_id,
)
from sglang.srt.distributed.parallel_state import GroupCoordinator, graph_capture
from sglang.srt.distributed.parallel_state import (
GroupCoordinator,
graph_capture,
set_pdmux_status,
)
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
get_attention_tp_rank,
@@ -53,6 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors,
enable_num_token_non_padded,
)
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.two_batch_overlap import TboCudaGraphRunnerPlugin
from sglang.srt.utils import (
empty_context,
@@ -249,6 +254,7 @@ class CudaGraphRunner:
self.tp_size = model_runner.server_args.tp_size
self.dp_size = model_runner.server_args.dp_size
self.pp_size = model_runner.server_args.pp_size
self.enable_pdmux = model_runner.server_args.enable_pdmux
self.attn_tp_size = get_attention_tp_size()
self.attn_tp_rank = get_attention_tp_rank()
@@ -286,6 +292,9 @@ class CudaGraphRunner:
self.model_runner.attn_backend.init_cuda_graph_state(
self.max_bs, self.max_num_token
)
# Init PDMux if needed
self.maybe_init_pdmux()
self.seq_len_fill_value = (
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
)
@@ -384,6 +393,12 @@ class CudaGraphRunner:
f"Capture cuda graph failed: {e}\n{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
)
def maybe_init_pdmux(self):
if self.enable_pdmux:
self.stream_groups = get_stream_groups()
for attn_backend in self.model_runner.decode_attn_backend_group:
attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
def _cache_loc_dtype(self):
return torch.int64
@@ -397,8 +412,12 @@ class CudaGraphRunner:
else:
cuda_graph_bs = forward_batch.batch_size
graph_key = cuda_graph_bs
if self.enable_pdmux:
graph_key = f"{get_current_stream_idx()}_{cuda_graph_bs}"
is_bs_supported = (
cuda_graph_bs in self.graphs
graph_key in self.graphs
if self.disable_padding
else cuda_graph_bs <= self.max_bs
)
@@ -478,52 +497,64 @@ class CudaGraphRunner:
if self.enable_profile_cuda_graph:
profile_context = self._init_profile_context_and_memory_record()
# Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes
# can reuse the memory pool allocated for the large shapes.
with freeze_gc(
self.model_runner.server_args.enable_cudagraph_gc
), graph_capture() as graph_capture_context:
with profile_context as prof:
self.stream = graph_capture_context.stream
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
empty_cache=False,
)
# Reverse the order to enable better memory sharing across cuda graphs.
capture_range = (
tqdm.tqdm(list(reversed(self.capture_bs)))
if get_tensor_model_parallel_rank() == 0
else reversed(self.capture_bs)
)
for i, bs in enumerate(capture_range):
if get_tensor_model_parallel_rank() == 0:
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
empty_cache=False,
)
capture_range.set_description(
f"Capturing batches ({bs=} {avail_mem=:.2f} GB)"
)
def _capture_one_stream(stream_idx: Optional[int] = None):
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
empty_cache=False,
)
# Reverse the order to enable better memory sharing across cuda graphs.
capture_range = (
tqdm.tqdm(list(reversed(self.capture_bs)))
if get_tensor_model_parallel_rank() == 0
else reversed(self.capture_bs)
)
for i, bs in enumerate(capture_range):
if get_tensor_model_parallel_rank() == 0:
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
empty_cache=False,
)
capture_range.set_description(
f"Capturing batches ({bs=} {avail_mem=:.2f} GB)"
)
with patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs,
tp_group=self.model_runner.tp_group,
) as forward:
(
graph,
output_buffers,
) = self.capture_one_batch_size(bs, forward)
self.graphs[bs] = graph
self.output_buffers[bs] = output_buffers
with patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs,
tp_group=self.model_runner.tp_group,
) as forward:
(
graph,
output_buffers,
) = self.capture_one_batch_size(bs, forward, stream_idx)
# For pd_multiplexing, we need to save the graph and output buffers
key = bs if stream_idx is None else f"{stream_idx}_{bs}"
self.graphs[key] = graph
self.output_buffers[key] = output_buffers
# Save gemlite cache after each capture
save_gemlite_cache()
# Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes
# can reuse the memory pool allocated for the large shapes.
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc):
if not self.enable_pdmux:
with graph_capture() as graph_capture_context, profile_context as prof:
self.stream = graph_capture_context.stream
_capture_one_stream()
else:
set_pdmux_status(False)
for i, sg in enumerate(self.stream_groups):
with graph_capture(
stream=sg[1]
) as graph_capture_context, profile_context as prof:
self.stream = graph_capture_context.stream
_capture_one_stream(i)
if self.enable_profile_cuda_graph:
self._post_process_after_profile(prof)
@@ -544,7 +575,9 @@ class CudaGraphRunner:
def _create_device_graph(self):
return torch.cuda.CUDAGraph()
def capture_one_batch_size(self, bs: int, forward: Callable):
def capture_one_batch_size(
self, bs: int, forward: Callable, stream_idx: Optional[int] = None
):
graph = self._create_device_graph()
stream = self.stream
num_tokens = bs * self.num_tokens_per_bs
@@ -618,6 +651,12 @@ class CudaGraphRunner:
else:
lora_ids = None
if stream_idx is None:
attn_backend = self.model_runner.attn_backend
else:
assert self.enable_pdmux
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
forward_batch = ForwardBatch(
forward_mode=self.capture_forward_mode,
batch_size=bs,
@@ -629,7 +668,7 @@ class CudaGraphRunner:
orig_seq_lens=seq_lens,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
attn_backend=attn_backend,
out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(),
encoder_lens=encoder_lens,
@@ -653,7 +692,7 @@ class CudaGraphRunner:
self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
# Attention backend
self.model_runner.attn_backend.init_forward_metadata_capture_cuda_graph(
attn_backend.init_forward_metadata_capture_cuda_graph(
bs,
num_tokens,
req_pool_indices,
@@ -814,7 +853,12 @@ class CudaGraphRunner:
if forward_batch.forward_mode.is_idle() and forward_batch.spec_info is not None:
forward_batch.spec_info.custom_mask = self.custom_mask
# Attention backend
self.model_runner.attn_backend.init_forward_metadata_replay_cuda_graph(
if self.enable_pdmux:
stream_idx = get_current_stream_idx()
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
else:
attn_backend = self.model_runner.attn_backend
attn_backend.init_forward_metadata_replay_cuda_graph(
bs,
self.req_pool_indices[:bs],
self.seq_lens[:bs],
@@ -846,9 +890,12 @@ class CudaGraphRunner:
self.positions[: self.raw_num_token].copy_(forward_batch.positions)
# Replay
self.graphs[self.bs].replay()
output = self.output_buffers[self.bs]
if self.enable_pdmux:
graph_key = f"{get_current_stream_idx()}_{self.bs}"
else:
graph_key = self.bs
self.graphs[graph_key].replay()
output = self.output_buffers[graph_key]
if isinstance(output, LogitsProcessorOutput):
return LogitsProcessorOutput(
next_token_logits=output.next_token_logits[: self.raw_num_token],

View File

@@ -63,6 +63,7 @@ class EAGLEDraftCudaGraphRunner:
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
self.enable_pdmux = False
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
server_args = model_runner.server_args
@@ -160,7 +161,9 @@ class EAGLEDraftCudaGraphRunner:
def capture(self):
CudaGraphRunner.capture(self)
def capture_one_batch_size(self, num_seqs: int, forward: Callable):
def capture_one_batch_size(
self, num_seqs: int, forward: Callable, stream_idx: int = 0
):
graph = torch.cuda.CUDAGraph()
stream = self.stream
num_tokens = num_seqs * self.num_tokens_per_bs

View File

@@ -61,6 +61,7 @@ class EAGLEDraftExtendCudaGraphRunner:
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
self.enable_pdmux = False
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(model_runner)
self.padded_static_len = -1
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
@@ -189,7 +190,7 @@ class EAGLEDraftExtendCudaGraphRunner:
def capture(self):
CudaGraphRunner.capture(self)
def capture_one_batch_size(self, bs: int, forward: Callable):
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
graph = torch.cuda.CUDAGraph()
stream = self.stream
num_tokens = bs * self.num_tokens_per_bs