Gather static input buffers for cuda graph (#13676)

This commit is contained in:
cctry
2025-11-22 13:01:50 -08:00
committed by GitHub
parent b29769f3b6
commit cad7878964
2 changed files with 245 additions and 145 deletions

View File

@@ -58,6 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors,
enable_num_token_non_padded,
)
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
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 (
@@ -302,9 +303,6 @@ class CudaGraphRunner:
)
self.encoder_len_fill_value = 0
self.seq_lens_cpu = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
)
if self.enable_torch_compile:
set_torch_compile_config()
@@ -315,83 +313,29 @@ class CudaGraphRunner:
num_tokens_per_bs=self.num_tokens_per_bs,
)
# Graph inputs
with torch.device(self.device):
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.input_embeds = torch.zeros(
(self.max_num_token, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.model_config.dtype,
)
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
self.seq_lens = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
)
self.out_cache_loc = torch.zeros(
(self.max_num_token,), dtype=self._cache_loc_dtype()
)
self.positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.mrope_positions = torch.zeros(
(3, self.max_num_token), dtype=torch.int64
)
self.num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
self.tbo_plugin = TboCudaGraphRunnerPlugin()
if self.require_gathered_buffer:
assert self.require_mlp_tp_gather or self.require_attn_tp_gather
self.buffers: GraphInputBuffers = GraphInputBuffers.create(
device=self.device,
max_bs=self.max_bs,
max_num_token=self.max_num_token,
hidden_size=self.model_runner.model_config.hidden_size,
vocab_size=self.model_runner.model_config.vocab_size,
dtype=self.model_runner.model_config.dtype,
dp_size=self.dp_size,
pp_size=self.pp_size,
is_encoder_decoder=self.is_encoder_decoder,
require_mlp_tp_gather=self.require_mlp_tp_gather,
seq_len_fill_value=self.seq_len_fill_value,
encoder_len_fill_value=self.encoder_len_fill_value,
num_tokens_per_bs=self.num_tokens_per_bs,
)
# pipeline parallelism
if self.pp_size > 1:
self.pp_proxy_tensors = {
"hidden_states": torch.zeros(
(self.max_bs, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.model_config.dtype,
),
"residual": torch.zeros(
(self.max_bs, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.model_config.dtype,
),
}
self.tbo_plugin = TboCudaGraphRunnerPlugin()
# Speculative_inference
if model_runner.spec_algorithm.is_eagle3():
self.model_runner.model.set_eagle3_layers_to_capture()
if self.is_encoder_decoder:
# NOTE: encoder_lens can influence the full_text_row_masked_out_mask tensor when doing mixed batch
self.encoder_lens = torch.full(
(self.max_bs,), self.encoder_len_fill_value, dtype=torch.int32
)
else:
self.encoder_lens = None
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
self.global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
)
self.global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
)
else:
assert self.require_attn_tp_gather
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
self.global_num_tokens_for_logprob_gpu = torch.zeros(
(1,), dtype=torch.int32
)
else:
self.global_num_tokens_gpu = None
self.global_num_tokens_for_logprob_gpu = None
self.custom_mask = torch.ones(
(
(self.seq_lens.sum().item() + self.max_num_token)
* self.num_tokens_per_bs
),
dtype=torch.bool,
device=self.device,
)
self.next_token_logits_buffer = torch.zeros(
(self.max_num_token, self.model_runner.model_config.vocab_size),
dtype=torch.float,
device=self.device,
)
# Speculative_inference
if model_runner.spec_algorithm.is_eagle3():
self.model_runner.model.set_eagle3_layers_to_capture()
# Capture
try:
@@ -587,40 +531,41 @@ class CudaGraphRunner:
def capture_one_batch_size(
self, bs: int, forward: Callable, stream_idx: Optional[int] = None
):
buffers = self.buffers
graph = self._create_device_graph()
stream = self.stream
num_tokens = bs * self.num_tokens_per_bs
# Graph inputs
input_ids = self.input_ids[:num_tokens]
req_pool_indices = self.req_pool_indices[:bs]
seq_lens = self.seq_lens[:bs]
seq_lens_cpu = self.seq_lens_cpu[:bs]
out_cache_loc = self.out_cache_loc[:num_tokens]
positions = self.positions[:num_tokens]
input_ids = buffers.input_ids[:num_tokens]
req_pool_indices = buffers.req_pool_indices[:bs]
seq_lens = buffers.seq_lens[:bs]
seq_lens_cpu = buffers.seq_lens_cpu[:bs]
out_cache_loc = buffers.out_cache_loc[:num_tokens]
positions = buffers.positions[:num_tokens]
if self.is_encoder_decoder:
encoder_lens = self.encoder_lens[:bs]
encoder_lens = buffers.encoder_lens[:bs]
else:
encoder_lens = None
mrope_positions = self.mrope_positions[:, :num_tokens]
next_token_logits_buffer = self.next_token_logits_buffer[:num_tokens]
self.num_token_non_padded[...] = num_tokens
mrope_positions = buffers.mrope_positions[:, :num_tokens]
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
buffers.num_token_non_padded[...] = num_tokens
# pipeline parallelism
if self.pp_size > 1:
pp_proxy_tensors = PPProxyTensors(
{k: v[:num_tokens] for k, v in self.pp_proxy_tensors.items()}
{k: v[:num_tokens] for k, v in buffers.pp_proxy_tensors.items()}
)
if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_(
buffers.global_num_tokens_gpu.copy_(
torch.tensor(
[num_tokens] * self.dp_size,
dtype=torch.int32,
device=input_ids.device,
)
)
self.global_num_tokens_for_logprob_gpu.copy_(
buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor(
[num_tokens] * self.dp_size,
dtype=torch.int32,
@@ -629,14 +574,14 @@ class CudaGraphRunner:
)
global_dp_buffer_len = num_tokens * self.dp_size
elif self.require_attn_tp_gather:
self.global_num_tokens_gpu.copy_(
buffers.global_num_tokens_gpu.copy_(
torch.tensor(
[num_tokens],
dtype=torch.int32,
device=input_ids.device,
)
)
self.global_num_tokens_for_logprob_gpu.copy_(
buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor(
[num_tokens],
dtype=torch.int32,
@@ -683,15 +628,15 @@ class CudaGraphRunner:
encoder_lens=encoder_lens,
return_logprob=False,
positions=positions,
global_num_tokens_gpu=self.global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu,
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len,
mrope_positions=mrope_positions,
spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info,
capture_hidden_mode=self.capture_hidden_mode,
num_token_non_padded=self.num_token_non_padded,
num_token_non_padded=buffers.num_token_non_padded,
global_forward_mode=self.capture_forward_mode,
lora_ids=lora_ids,
)
@@ -793,6 +738,7 @@ class CudaGraphRunner:
forward_batch: ForwardBatch,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
):
buffers = self.buffers
self.recapture_if_needed(forward_batch)
raw_bs = forward_batch.batch_size
@@ -810,48 +756,23 @@ class CudaGraphRunner:
else:
index = bisect.bisect_left(self.capture_bs, raw_bs)
bs = self.capture_bs[index]
if bs != raw_bs:
self.seq_lens.fill_(self.seq_len_fill_value)
self.out_cache_loc.zero_()
# Common inputs
self.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
self.out_cache_loc[:raw_num_token].copy_(forward_batch.out_cache_loc)
self.positions[:raw_num_token].copy_(forward_batch.positions)
seq_lens_cpu = None
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
self.seq_lens_cpu.fill_(self.seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
seq_lens_cpu = self.seq_lens_cpu[:bs]
if pp_proxy_tensors:
for key in self.pp_proxy_tensors.keys():
dim = pp_proxy_tensors[key].shape[0]
self.pp_proxy_tensors[key][:dim].copy_(pp_proxy_tensors[key])
if self.is_encoder_decoder:
self.encoder_lens[:raw_bs].copy_(forward_batch.encoder_lens)
if forward_batch.mrope_positions is not None:
self.mrope_positions[:, :raw_num_token].copy_(forward_batch.mrope_positions)
if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
self.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
if enable_num_token_non_padded(self.model_runner.server_args):
num_token_non_padded = forward_batch.num_token_non_padded
if self.require_gathered_buffer and not self.nsa_enable_prefill_cp:
tokens_per_rank = bs // self.attn_tp_size * self.num_tokens_per_bs
num_local_token_non_padded = torch.clamp(
num_token_non_padded - tokens_per_rank * self.attn_tp_rank,
min=0,
max=tokens_per_rank,
)
self.num_token_non_padded.copy_(num_local_token_non_padded)
else:
self.num_token_non_padded.copy_(num_token_non_padded)
seq_lens_cpu = buffers.populate_from_forward_batch(
forward_batch=forward_batch,
raw_bs=raw_bs,
raw_num_token=raw_num_token,
bs=bs,
seq_len_fill_value=self.seq_len_fill_value,
require_gathered_buffer=self.require_gathered_buffer,
num_tokens_per_bs=self.num_tokens_per_bs,
nsa_enable_prefill_cp=self.nsa_enable_prefill_cp,
attn_tp_rank=self.attn_tp_rank,
attn_tp_size=self.attn_tp_size,
enable_num_token_non_padded_flag=enable_num_token_non_padded(
self.model_runner.server_args
),
pp_proxy_tensors=pp_proxy_tensors,
)
if self.enable_two_batch_overlap:
self.tbo_plugin.replay_prepare(
forward_mode=self.capture_forward_mode,
@@ -860,7 +781,7 @@ class CudaGraphRunner:
spec_info=forward_batch.spec_info,
)
if forward_batch.forward_mode.is_idle() and forward_batch.spec_info is not None:
forward_batch.spec_info.custom_mask = self.custom_mask
forward_batch.spec_info.custom_mask = buffers.custom_mask
# Attention backend
if self.enable_pdmux:
stream_idx = get_current_stream_idx()
@@ -869,10 +790,10 @@ class CudaGraphRunner:
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],
buffers.req_pool_indices[:bs],
buffers.seq_lens[:bs],
forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value,
self.encoder_lens[:bs] if self.is_encoder_decoder else None,
buffers.encoder_lens[:bs] if self.is_encoder_decoder else None,
self.capture_forward_mode,
forward_batch.spec_info,
seq_lens_cpu=seq_lens_cpu,
@@ -895,8 +816,8 @@ class CudaGraphRunner:
self.replay_prepare(forward_batch, pp_proxy_tensors)
else:
# In speculative decoding, these two fields are still needed.
self.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
self.positions[: self.raw_num_token].copy_(forward_batch.positions)
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
# Replay
if self.enable_pdmux:
@@ -931,7 +852,7 @@ class CudaGraphRunner:
else:
spec_info = EagleVerifyInput(
draft_token=None,
custom_mask=self.custom_mask,
custom_mask=self.buffers.custom_mask,
positions=None,
retrive_index=None,
retrive_next_token=None,
@@ -950,7 +871,7 @@ class CudaGraphRunner:
spec_info = NgramVerifyInput(
draft_token=None,
tree_mask=self.custom_mask,
tree_mask=self.buffers.custom_mask,
positions=None,
retrive_index=None,
retrive_next_token=None,

View File

@@ -0,0 +1,179 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Optional
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
@dataclass
class GraphInputBuffers:
input_ids: torch.Tensor
input_embeds: torch.Tensor
req_pool_indices: torch.Tensor
seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor
out_cache_loc: torch.Tensor
positions: torch.Tensor
mrope_positions: torch.Tensor
num_token_non_padded: torch.Tensor
custom_mask: torch.Tensor
next_token_logits_buffer: torch.Tensor
global_num_tokens_gpu: torch.Tensor
global_num_tokens_for_logprob_gpu: torch.Tensor
encoder_lens: Optional[torch.Tensor]
pp_proxy_tensors: Optional[Dict[str, torch.Tensor]]
@classmethod
def create(
cls,
*,
device: torch.device,
max_bs: int,
max_num_token: int,
hidden_size: int,
vocab_size: int,
dtype: torch.dtype,
dp_size: int,
pp_size: int,
is_encoder_decoder: bool,
require_mlp_tp_gather: bool,
seq_len_fill_value: int,
encoder_len_fill_value: int,
num_tokens_per_bs: int,
) -> "GraphInputBuffers":
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int32)
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int32)
out_cache_loc = torch.zeros((max_num_token,), dtype=torch.int64)
positions = torch.zeros((max_num_token,), dtype=torch.int64)
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
custom_mask = torch.ones(
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
dtype=torch.bool,
)
next_token_logits_buffer = torch.zeros(
(max_num_token, vocab_size),
dtype=torch.float,
)
if pp_size > 1:
pp_proxy_tensors = {
"hidden_states": torch.zeros((max_bs, hidden_size), dtype=dtype),
"residual": torch.zeros((max_bs, hidden_size), dtype=dtype),
}
else:
pp_proxy_tensors = None
if is_encoder_decoder:
encoder_lens = torch.full(
(max_bs,), encoder_len_fill_value, dtype=torch.int32
)
else:
encoder_lens = None
if require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros(
(dp_size,), dtype=torch.int32
)
else:
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
# Keep seq_lens_cpu as a true CPU tensor, like the old implementation.
seq_lens_cpu = torch.full(
(max_bs,),
seq_len_fill_value,
dtype=torch.int32,
device="cpu",
)
return cls(
input_ids=input_ids,
input_embeds=input_embeds,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
positions=positions,
mrope_positions=mrope_positions,
num_token_non_padded=num_token_non_padded,
custom_mask=custom_mask,
next_token_logits_buffer=next_token_logits_buffer,
encoder_lens=encoder_lens,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
pp_proxy_tensors=pp_proxy_tensors,
)
def populate_from_forward_batch(
self,
*,
forward_batch: ForwardBatch,
raw_bs: int,
raw_num_token: int,
bs: int,
seq_len_fill_value: int,
require_gathered_buffer: bool,
num_tokens_per_bs: int,
nsa_enable_prefill_cp: bool,
attn_tp_rank: int,
attn_tp_size: int,
enable_num_token_non_padded_flag: bool,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> Optional[torch.Tensor]:
if bs != raw_bs:
self.seq_lens.fill_(seq_len_fill_value)
self.out_cache_loc.zero_()
# Common inputs
self.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
self.out_cache_loc[:raw_num_token].copy_(forward_batch.out_cache_loc)
self.positions[:raw_num_token].copy_(forward_batch.positions)
seq_lens_cpu: Optional[torch.Tensor] = None
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
self.seq_lens_cpu.fill_(seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
seq_lens_cpu = self.seq_lens_cpu[:bs]
if self.encoder_lens is not None and forward_batch.encoder_lens is not None:
self.encoder_lens[:raw_bs].copy_(forward_batch.encoder_lens)
if forward_batch.mrope_positions is not None:
self.mrope_positions[:, :raw_num_token].copy_(forward_batch.mrope_positions)
if require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_bs)
if enable_num_token_non_padded_flag:
num_token_non_padded = forward_batch.num_token_non_padded
if require_gathered_buffer and not nsa_enable_prefill_cp:
tokens_per_rank = bs // attn_tp_size * num_tokens_per_bs
num_local_token_non_padded = torch.clamp(
num_token_non_padded - tokens_per_rank * attn_tp_rank,
min=0,
max=tokens_per_rank,
)
self.num_token_non_padded.copy_(num_local_token_non_padded)
else:
self.num_token_non_padded.copy_(num_token_non_padded)
# Pipeline-parallel proxy tensors.
if pp_proxy_tensors is not None and self.pp_proxy_tensors is not None:
for key, buf in self.pp_proxy_tensors.items():
src = pp_proxy_tensors.tensors[key]
dim = src.shape[0]
buf[:dim].copy_(src)
return seq_lens_cpu