diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index a95a4400c..2fb40cdec 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -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, diff --git a/python/sglang/srt/model_executor/input_buffers.py b/python/sglang/srt/model_executor/input_buffers.py new file mode 100644 index 000000000..586399474 --- /dev/null +++ b/python/sglang/srt/model_executor/input_buffers.py @@ -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