Refactor CUDA graph input buffers with shared buffer pool (#19180)

This commit is contained in:
Cheng Wan
2026-02-24 02:24:40 -08:00
committed by GitHub
parent 94946764a4
commit 6e54361608
8 changed files with 699 additions and 494 deletions
@@ -21,8 +21,9 @@ import inspect
import logging import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass
from functools import partial from functools import partial
from typing import TYPE_CHECKING, Callable, Optional, Union from typing import TYPE_CHECKING, Callable, Dict, Optional, Union
import torch import torch
import tqdm import tqdm
@@ -58,9 +59,10 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
compute_local_num_token_non_padded,
enable_num_token_non_padded, enable_num_token_non_padded,
) )
from sglang.srt.model_executor.input_buffers import GraphInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
@@ -90,6 +92,200 @@ logger = logging.getLogger(__name__)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@dataclass
class DecodeInputBuffers(ForwardInputBuffers):
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
mamba_track_indices: Optional[torch.Tensor]
mamba_track_mask: Optional[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,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
) -> "DecodeInputBuffers":
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=cache_loc_dtype)
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,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64)
if enable_mamba_track
else None
)
mamba_track_mask = (
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
)
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,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
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,
enable_num_token_non_padded_flag: bool,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
):
if bs != raw_bs:
self.seq_lens.fill_(seq_len_fill_value)
self.out_cache_loc.zero_()
if self.mamba_track_indices is not None:
self.mamba_track_indices.zero_()
if self.mamba_track_mask is not None:
self.mamba_track_mask.fill_(False)
# 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)
if (
self.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None
):
self.mamba_track_indices[:raw_bs].copy_(forward_batch.mamba_track_indices)
if (
self.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None
):
self.mamba_track_mask[:raw_bs].copy_(forward_batch.mamba_track_mask)
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)
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:
if require_gathered_buffer and not nsa_enable_prefill_cp:
num_tokens_per_dp = bs * num_tokens_per_bs
local = compute_local_num_token_non_padded(
global_num_token_non_padded=forward_batch.num_token_non_padded,
num_tokens_per_dp=num_tokens_per_dp,
)
self.num_token_non_padded.copy_(local)
else:
self.num_token_non_padded.copy_(forward_batch.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)
# Detect whether the current forward pass is in capture mode # Detect whether the current forward pass is in capture mode
is_capture_mode = False is_capture_mode = False
@@ -337,7 +533,7 @@ class CudaGraphRunner:
if self.require_gathered_buffer: if self.require_gathered_buffer:
assert self.require_mlp_tp_gather or self.require_attn_tp_gather assert self.require_mlp_tp_gather or self.require_attn_tp_gather
self.buffers: GraphInputBuffers = GraphInputBuffers.create( self.buffers: DecodeInputBuffers = DecodeInputBuffers.create(
device=self.device, device=self.device,
max_bs=self.max_bs, max_bs=self.max_bs,
max_num_token=self.max_num_token, max_num_token=self.max_num_token,
@@ -354,6 +550,7 @@ class CudaGraphRunner:
cache_loc_dtype=self._cache_loc_dtype(), cache_loc_dtype=self._cache_loc_dtype(),
enable_mamba_track=enable_mamba_track, enable_mamba_track=enable_mamba_track,
) )
self.buffers.share_buffers()
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
@@ -556,7 +753,7 @@ class CudaGraphRunner:
def capture_one_batch_size( def capture_one_batch_size(
self, bs: int, forward: Callable, stream_idx: Optional[int] = None self, bs: int, forward: Callable, stream_idx: Optional[int] = None
): ):
buffers: GraphInputBuffers = self.buffers buffers: DecodeInputBuffers = self.buffers
graph = self._create_device_graph() graph = self._create_device_graph()
stream = self.stream stream = self.stream
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
@@ -798,7 +995,7 @@ class CudaGraphRunner:
index = bisect.bisect_left(self.capture_bs, raw_bs) index = bisect.bisect_left(self.capture_bs, raw_bs)
bs = self.capture_bs[index] bs = self.capture_bs[index]
seq_lens_cpu = buffers.populate_from_forward_batch( buffers.populate_from_forward_batch(
forward_batch=forward_batch, forward_batch=forward_batch,
raw_bs=raw_bs, raw_bs=raw_bs,
raw_num_token=raw_num_token, raw_num_token=raw_num_token,
@@ -835,7 +1032,7 @@ class CudaGraphRunner:
buffers.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, self.capture_forward_mode,
forward_batch.spec_info, forward_batch.spec_info,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=buffers.seq_lens_cpu[:bs],
) )
# Store fields # Store fields
+44 -197
View File
@@ -1,208 +1,55 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass, fields
from typing import Dict, Optional from typing import Dict
import torch import torch
from sglang.srt.model_executor.forward_batch_info import ( _forward_input_buffer_pool: Dict[str, torch.Tensor] = {}
ForwardBatch,
PPProxyTensors,
compute_local_num_token_non_padded,
)
@dataclass @dataclass
class GraphInputBuffers: class ForwardInputBuffers:
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
mamba_track_indices: Optional[torch.Tensor]
mamba_track_mask: Optional[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 _share_one_buffer(self, name: str, new_buffer: torch.Tensor) -> torch.Tensor:
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,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
) -> "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=cache_loc_dtype)
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,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64)
if enable_mamba_track
else None
)
mamba_track_mask = (
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
)
if pp_size > 1: buffer_size = new_buffer.size()
pp_proxy_tensors = { buffer_stride = new_buffer.stride()
"hidden_states": torch.zeros((max_bs, hidden_size), dtype=dtype),
"residual": torch.zeros((max_bs, hidden_size), dtype=dtype), old_buffer = _forward_input_buffer_pool.get(name, None)
} if old_buffer is not None:
assert (
new_buffer.dtype == old_buffer.dtype
), f"Buffer {name} has different dtype than before."
assert (
new_buffer.device == old_buffer.device
), f"Buffer {name} has different device than before."
if old_buffer.numel() > new_buffer.numel():
new_buffer = old_buffer
_forward_input_buffer_pool[name] = new_buffer
return new_buffer.as_strided(buffer_size, buffer_stride)
def share_buffers(self):
for f in fields(self):
name = f.name
buffer = getattr(self, name)
if buffer is None:
continue
elif isinstance(buffer, dict):
for sub_name, sub_buffer in buffer.items():
assert isinstance(
sub_buffer, torch.Tensor
), f"Field {name}.{sub_name} is expected to be a torch.Tensor, but got {type(sub_buffer)}."
new_buffer = self._share_one_buffer(
f"{name}.{sub_name}", sub_buffer
)
buffer[sub_name] = new_buffer
else: else:
pp_proxy_tensors = None assert isinstance(
buffer, torch.Tensor
if is_encoder_decoder: ), f"Field {name} is expected to be a torch.Tensor or a dict of torch.Tensor, but got {type(buffer)}."
encoder_lens = torch.full( new_buffer = self._share_one_buffer(name, buffer)
(max_bs,), encoder_len_fill_value, dtype=torch.int32 setattr(self, name, new_buffer)
)
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,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
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,
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_()
if self.mamba_track_indices is not None:
self.mamba_track_indices.zero_()
if self.mamba_track_mask is not None:
self.mamba_track_mask.fill_(False)
# 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)
if (
self.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None
):
self.mamba_track_indices[:raw_bs].copy_(forward_batch.mamba_track_indices)
if (
self.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None
):
self.mamba_track_mask[:raw_bs].copy_(forward_batch.mamba_track_mask)
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:
if require_gathered_buffer and not nsa_enable_prefill_cp:
num_tokens_per_dp = bs * num_tokens_per_bs
local = compute_local_num_token_non_padded(
global_num_token_non_padded=forward_batch.num_token_non_padded,
num_tokens_per_dp=num_tokens_per_dp,
)
self.num_token_non_padded.copy_(local)
else:
self.num_token_non_padded.copy_(forward_batch.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
@@ -118,6 +118,7 @@ from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
from sglang.srt.model_executor.cuda_graph_runner import ( from sglang.srt.model_executor.cuda_graph_runner import (
CudaGraphRunner, CudaGraphRunner,
DecodeInputBuffers,
set_torch_compile_config, set_torch_compile_config,
) )
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
@@ -127,7 +128,6 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
from sglang.srt.model_executor.model_runner_kv_cache_mixin import ( from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
ModelRunnerKVCacheMixin, ModelRunnerKVCacheMixin,
) )
@@ -1908,7 +1908,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if require_gathered_buffer(self.server_args): if require_gathered_buffer(self.server_args):
assert require_mlp_tp_gather_ or require_attn_tp_gather(self.server_args) assert require_mlp_tp_gather_ or require_attn_tp_gather(self.server_args)
buffers: GraphInputBuffers = GraphInputBuffers.create( buffers: DecodeInputBuffers = DecodeInputBuffers.create(
device=self.device, device=self.device,
max_bs=batch_size, max_bs=batch_size,
max_num_token=num_tokens, max_num_token=num_tokens,
@@ -19,7 +19,8 @@ import bisect
import gc import gc
import logging import logging
from contextlib import contextmanager from contextlib import contextmanager
from typing import TYPE_CHECKING, Union from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional, Union
import torch import torch
import tqdm import tqdm
@@ -55,6 +56,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.utils import get_available_gpu_memory, is_npu, log_info_on_rank0 from sglang.srt.utils import get_available_gpu_memory, is_npu, log_info_on_rank0
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -63,6 +65,19 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@dataclass
class PrefillInputBuffers(ForwardInputBuffers):
input_ids: torch.Tensor
out_cache_loc: torch.Tensor
out_cache_loc_swa: Optional[torch.Tensor]
mamba_track_indices: Optional[torch.Tensor]
mamba_track_mask: Optional[torch.Tensor]
mamba_track_seqlens: Optional[torch.Tensor]
positions: torch.Tensor
input_embeds: Optional[torch.Tensor]
mrope_positions: Optional[torch.Tensor]
@contextmanager @contextmanager
def freeze_gc(enable_cudagraph_gc: bool): def freeze_gc(enable_cudagraph_gc: bool):
""" """
@@ -189,31 +204,31 @@ class PiecewiseCudaGraphRunner:
# Graph inputs # Graph inputs
with torch.device(self.device): with torch.device(self.device):
self.input_ids = torch.zeros((self.max_num_tokens,), dtype=torch.int64) input_ids = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
self.out_cache_loc = torch.zeros( out_cache_loc = torch.zeros(
(self.max_num_tokens,), dtype=self._cache_loc_dtype() (self.max_num_tokens,), dtype=self._cache_loc_dtype()
) )
self.out_cache_loc_swa = ( out_cache_loc_swa = (
torch.zeros((self.max_num_tokens,), dtype=torch.int64) torch.zeros((self.max_num_tokens,), dtype=torch.int64)
if model_runner.is_hybrid_swa if model_runner.is_hybrid_swa
else None else None
) )
self.mamba_track_indices = ( mamba_track_indices = (
torch.zeros((self.max_bs,), dtype=torch.int64) torch.zeros((self.max_bs,), dtype=torch.int64)
if self.mamba_track_enabled if self.mamba_track_enabled
else None else None
) )
self.mamba_track_mask = ( mamba_track_mask = (
torch.zeros((self.max_bs,), dtype=torch.bool) torch.zeros((self.max_bs,), dtype=torch.bool)
if self.mamba_track_enabled if self.mamba_track_enabled
else None else None
) )
self.mamba_track_seqlens = ( mamba_track_seqlens = (
torch.zeros((self.max_bs,), dtype=torch.int32) torch.zeros((self.max_bs,), dtype=torch.int32)
if self.mamba_track_enabled if self.mamba_track_enabled
else None else None
) )
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64) positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
@@ -223,13 +238,29 @@ class PiecewiseCudaGraphRunner:
# 1. In multimodal, we only compile and capture the language model part. # 1. In multimodal, we only compile and capture the language model part.
# 2. The embedder is outside of the graph, but cuda graph requires the input embeds to have a fixed memory address. # 2. The embedder is outside of the graph, but cuda graph requires the input embeds to have a fixed memory address.
# 3. Input embeds is a pre-allocated buffer. In model.forward, we copy the embed output to this buffer. # 3. Input embeds is a pre-allocated buffer. In model.forward, we copy the embed output to this buffer.
self.input_embeds = torch.zeros( input_embeds = torch.zeros(
(self.max_num_tokens, self.model_runner.model_config.hidden_size), (self.max_num_tokens, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
self.mrope_positions = torch.zeros( mrope_positions = torch.zeros(
(3, self.max_num_tokens), dtype=torch.int64 (3, self.max_num_tokens), dtype=torch.int64
) )
else:
input_embeds = None
mrope_positions = None
self.buffers = PrefillInputBuffers(
input_ids=input_ids,
out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
mamba_track_seqlens=mamba_track_seqlens,
positions=positions,
input_embeds=input_embeds,
mrope_positions=mrope_positions,
)
self.buffers.share_buffers()
self.attention_layers = self.model_runner.attention_layers self.attention_layers = self.model_runner.attention_layers
self.moe_layers = self.model_runner.moe_layers self.moe_layers = self.model_runner.moe_layers
@@ -285,29 +316,32 @@ class PiecewiseCudaGraphRunner:
def warmup_torch_compile(self, num_tokens: int): def warmup_torch_compile(self, num_tokens: int):
"""Warmup the model with a simple forward pass before CUDA graph capture.""" """Warmup the model with a simple forward pass before CUDA graph capture."""
input_ids = self.input_ids[:num_tokens] buffers = self.buffers
input_embeds = self.input_embeds[:num_tokens] if self.is_multimodal else None input_ids = buffers.input_ids[:num_tokens]
positions = self.positions[:num_tokens] input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None
positions = buffers.positions[:num_tokens]
mrope_positions = ( mrope_positions = (
self.mrope_positions[:, :num_tokens] if self.is_multimodal else None buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
) )
out_cache_loc = self.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
out_cache_loc_swa = ( out_cache_loc_swa = (
self.out_cache_loc_swa[:num_tokens] buffers.out_cache_loc_swa[:num_tokens]
if self.out_cache_loc_swa is not None if buffers.out_cache_loc_swa is not None
else None else None
) )
mamba_track_indices = ( mamba_track_indices = (
self.mamba_track_indices[:1] buffers.mamba_track_indices[:1]
if self.mamba_track_indices is not None if buffers.mamba_track_indices is not None
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
self.mamba_track_mask[:1] if self.mamba_track_mask is not None else None buffers.mamba_track_mask[:1]
if buffers.mamba_track_mask is not None
else None
) )
mamba_track_seqlens = ( mamba_track_seqlens = (
self.mamba_track_seqlens[:1] buffers.mamba_track_seqlens[:1]
if self.mamba_track_seqlens is not None if buffers.mamba_track_seqlens is not None
else None else None
) )
with torch.device(self.device): with torch.device(self.device):
@@ -422,34 +456,37 @@ class PiecewiseCudaGraphRunner:
self.capture_one_batch_size(num_tokens) self.capture_one_batch_size(num_tokens)
def capture_one_batch_size(self, num_tokens: int): def capture_one_batch_size(self, num_tokens: int):
buffers = self.buffers
bs = 1 bs = 1
# Graph inputs # Graph inputs
input_ids = self.input_ids[:num_tokens] input_ids = buffers.input_ids[:num_tokens]
input_embeds = self.input_embeds[:num_tokens] if self.is_multimodal else None input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None
out_cache_loc = self.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
out_cache_loc_swa = ( out_cache_loc_swa = (
self.out_cache_loc_swa[:num_tokens] buffers.out_cache_loc_swa[:num_tokens]
if self.out_cache_loc_swa is not None if buffers.out_cache_loc_swa is not None
else None else None
) )
mamba_track_indices = ( mamba_track_indices = (
self.mamba_track_indices[:bs] buffers.mamba_track_indices[:bs]
if self.mamba_track_indices is not None if buffers.mamba_track_indices is not None
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None buffers.mamba_track_mask[:bs]
) if buffers.mamba_track_mask is not None
mamba_track_seqlens = (
self.mamba_track_seqlens[:bs]
if self.mamba_track_seqlens is not None
else None else None
) )
positions = self.positions[:num_tokens] mamba_track_seqlens = (
buffers.mamba_track_seqlens[:bs]
if buffers.mamba_track_seqlens is not None
else None
)
positions = buffers.positions[:num_tokens]
mrope_positions = ( mrope_positions = (
self.mrope_positions[:, :num_tokens] if self.is_multimodal else None buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
) )
global_dp_buffer_len = None global_dp_buffer_len = None
@@ -553,82 +590,85 @@ class PiecewiseCudaGraphRunner:
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
**kwargs, **kwargs,
): ):
buffers = self.buffers
num_tokens = len(forward_batch.input_ids) num_tokens = len(forward_batch.input_ids)
index = bisect.bisect_left(self.capture_num_tokens, num_tokens) index = bisect.bisect_left(self.capture_num_tokens, num_tokens)
static_num_tokens = self.capture_num_tokens[index] static_num_tokens = self.capture_num_tokens[index]
self.raw_num_tokens = num_tokens self.raw_num_tokens = num_tokens
if static_num_tokens != num_tokens: if static_num_tokens != num_tokens:
self.out_cache_loc.zero_() buffers.out_cache_loc.zero_()
if self.out_cache_loc_swa is not None: if buffers.out_cache_loc_swa is not None:
self.out_cache_loc_swa.zero_() buffers.out_cache_loc_swa.zero_()
self.input_ids[num_tokens:static_num_tokens].zero_() buffers.input_ids[num_tokens:static_num_tokens].zero_()
self.positions[num_tokens:static_num_tokens].zero_() buffers.positions[num_tokens:static_num_tokens].zero_()
if self.is_multimodal: if self.is_multimodal:
self.input_embeds[:, num_tokens:static_num_tokens].zero_() buffers.input_embeds[:, num_tokens:static_num_tokens].zero_()
if forward_batch.mrope_positions is not None: if forward_batch.mrope_positions is not None:
self.mrope_positions[:, num_tokens:static_num_tokens].zero_() buffers.mrope_positions[:, num_tokens:static_num_tokens].zero_()
bs = forward_batch.batch_size bs = forward_batch.batch_size
self.input_ids[:num_tokens].copy_(forward_batch.input_ids) buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
self.positions[:num_tokens].copy_(forward_batch.positions) buffers.positions[:num_tokens].copy_(forward_batch.positions)
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc) buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
if self.out_cache_loc_swa is not None: if buffers.out_cache_loc_swa is not None:
self.out_cache_loc_swa[: self.raw_num_tokens].copy_( buffers.out_cache_loc_swa[: self.raw_num_tokens].copy_(
self.model_runner.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( self.model_runner.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc forward_batch.out_cache_loc
) )
) )
if ( if (
self.mamba_track_indices is not None buffers.mamba_track_indices is not None
and forward_batch.mamba_track_indices is not None and forward_batch.mamba_track_indices is not None
): ):
self.mamba_track_indices[:bs].copy_(forward_batch.mamba_track_indices) buffers.mamba_track_indices[:bs].copy_(forward_batch.mamba_track_indices)
if ( if (
self.mamba_track_mask is not None buffers.mamba_track_mask is not None
and forward_batch.mamba_track_mask is not None and forward_batch.mamba_track_mask is not None
): ):
self.mamba_track_mask[:bs].copy_(forward_batch.mamba_track_mask) buffers.mamba_track_mask[:bs].copy_(forward_batch.mamba_track_mask)
if ( if (
self.mamba_track_seqlens is not None buffers.mamba_track_seqlens is not None
and forward_batch.mamba_track_seqlens is not None and forward_batch.mamba_track_seqlens is not None
): ):
self.mamba_track_seqlens[:bs].copy_(forward_batch.mamba_track_seqlens) buffers.mamba_track_seqlens[:bs].copy_(forward_batch.mamba_track_seqlens)
input_ids = self.input_ids[:static_num_tokens] input_ids = buffers.input_ids[:static_num_tokens]
positions = self.positions[:static_num_tokens] positions = buffers.positions[:static_num_tokens]
out_cache_loc = self.out_cache_loc[:static_num_tokens] out_cache_loc = buffers.out_cache_loc[:static_num_tokens]
out_cache_loc_swa = ( out_cache_loc_swa = (
self.out_cache_loc_swa[:static_num_tokens] buffers.out_cache_loc_swa[:static_num_tokens]
if forward_batch.out_cache_loc_swa is not None if forward_batch.out_cache_loc_swa is not None
else None else None
) )
mamba_track_indices = ( mamba_track_indices = (
self.mamba_track_indices[:bs] buffers.mamba_track_indices[:bs]
if self.mamba_track_indices is not None if buffers.mamba_track_indices is not None
else None else None
) )
mamba_track_mask = ( mamba_track_mask = (
self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None buffers.mamba_track_mask[:bs]
if buffers.mamba_track_mask is not None
else None
) )
mamba_track_seqlens = ( mamba_track_seqlens = (
self.mamba_track_seqlens[:bs] buffers.mamba_track_seqlens[:bs]
if self.mamba_track_seqlens is not None if buffers.mamba_track_seqlens is not None
else None else None
) )
if forward_batch.mrope_positions is not None: if forward_batch.mrope_positions is not None:
self.mrope_positions[:, :num_tokens].copy_(forward_batch.mrope_positions) buffers.mrope_positions[:, :num_tokens].copy_(forward_batch.mrope_positions)
input_ids = self.input_ids[:static_num_tokens] input_ids = buffers.input_ids[:static_num_tokens]
input_embeds = ( input_embeds = (
self.input_embeds[:static_num_tokens] if self.is_multimodal else None buffers.input_embeds[:static_num_tokens] if self.is_multimodal else None
) )
mrope_positions = ( mrope_positions = (
self.mrope_positions[:, :static_num_tokens] buffers.mrope_positions[:, :static_num_tokens]
if forward_batch.mrope_positions is not None if forward_batch.mrope_positions is not None
else None else None
) )
@@ -1,7 +1,8 @@
from __future__ import annotations from __future__ import annotations
import bisect import bisect
from typing import TYPE_CHECKING, Callable from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, Optional
import torch import torch
@@ -22,6 +23,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.utils import ( from sglang.srt.utils import (
require_attn_tp_gather, require_attn_tp_gather,
@@ -34,6 +36,23 @@ if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker import EAGLEWorker from sglang.srt.speculative.eagle_worker import EAGLEWorker
@dataclass
class EagleDraftInputBuffers(ForwardInputBuffers):
input_ids: torch.Tensor
req_pool_indices: torch.Tensor
out_cache_loc: torch.Tensor
positions: torch.Tensor
mrope_positions: torch.Tensor
seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor
extend_seq_lens: torch.Tensor
topk_p: torch.Tensor
topk_index: torch.Tensor
hidden_states: torch.Tensor
global_num_tokens_gpu: Optional[torch.Tensor]
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
class EAGLEDraftCudaGraphRunner: class EAGLEDraftCudaGraphRunner:
def __init__(self, eagle_worker: EAGLEWorker): def __init__(self, eagle_worker: EAGLEWorker):
# Parse args # Parse args
@@ -75,7 +94,7 @@ class EAGLEDraftCudaGraphRunner:
self.seq_len_fill_value = self.model_runner.draft_attn_backend.attn_backends[ self.seq_len_fill_value = self.model_runner.draft_attn_backend.attn_backends[
0 0
].get_cuda_graph_seq_len_fill_value() ].get_cuda_graph_seq_len_fill_value()
self.seq_lens_cpu = torch.full( seq_lens_cpu = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32 (self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
) )
self.extend_seq_lens_cpu = [self.seq_len_fill_value] * self.max_bs self.extend_seq_lens_cpu = [self.seq_len_fill_value] * self.max_bs
@@ -85,44 +104,59 @@ class EAGLEDraftCudaGraphRunner:
# Graph inputs # Graph inputs
with torch.device(model_runner.device): with torch.device(model_runner.device):
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64) input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32) req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
self.out_cache_loc = torch.zeros( out_cache_loc = torch.zeros(
(self.max_num_token * self.speculative_num_steps,), (self.max_num_token * self.speculative_num_steps,),
dtype=self._cache_loc_dtype(), dtype=self._cache_loc_dtype(),
) )
self.positions = torch.zeros((self.max_num_token,), dtype=torch.int64) positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.mrope_positions = torch.zeros( mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
(3, self.max_num_token), dtype=torch.int64 seq_lens = torch.full(
)
self.seq_lens = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32 (self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
) )
self.extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32) extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
self.topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32) topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32)
self.topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64) topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
self.hidden_states = torch.zeros( hidden_states = torch.zeros(
(self.max_bs, self.model_runner.model_config.hidden_size), (self.max_bs, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
if self.require_gathered_buffer: if self.require_gathered_buffer:
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu = torch.zeros( global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
else: else:
assert self.require_attn_tp_gather assert self.require_attn_tp_gather
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32) global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(1,), dtype=torch.int32 (1,), dtype=torch.int32
) )
else: else:
self.global_num_tokens_gpu = None global_num_tokens_gpu = None
self.global_num_tokens_for_logprob_gpu = None global_num_tokens_for_logprob_gpu = None
self.buffers = EagleDraftInputBuffers(
input_ids=input_ids,
req_pool_indices=req_pool_indices,
out_cache_loc=out_cache_loc,
positions=positions,
mrope_positions=mrope_positions,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
extend_seq_lens=extend_seq_lens,
topk_p=topk_p,
topk_index=topk_index,
hidden_states=hidden_states,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
)
self.buffers.share_buffers()
# Capture # Capture
try: try:
@@ -181,59 +215,60 @@ class EAGLEDraftCudaGraphRunner:
def capture_one_batch_size( def capture_one_batch_size(
self, num_seqs: int, forward: Callable, stream_idx: int = 0 self, num_seqs: int, forward: Callable, stream_idx: int = 0
): ):
buffers = self.buffers
graph = self._create_graph() graph = self._create_graph()
stream = self.stream stream = self.stream
num_tokens = num_seqs * self.num_tokens_per_bs num_tokens = num_seqs * self.num_tokens_per_bs
# Graph inputs # Graph inputs
req_pool_indices = self.req_pool_indices[:num_seqs] req_pool_indices = buffers.req_pool_indices[:num_seqs]
seq_lens = self.seq_lens[:num_seqs] seq_lens = buffers.seq_lens[:num_seqs]
seq_lens_cpu = self.seq_lens_cpu[:num_seqs] seq_lens_cpu = buffers.seq_lens_cpu[:num_seqs]
extend_seq_lens = self.extend_seq_lens[:num_seqs] extend_seq_lens = buffers.extend_seq_lens[:num_seqs]
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:num_seqs] extend_seq_lens_cpu = self.extend_seq_lens_cpu[:num_seqs]
out_cache_loc = self.out_cache_loc[: num_tokens * self.speculative_num_steps] out_cache_loc = buffers.out_cache_loc[: num_tokens * self.speculative_num_steps]
positions = self.positions[:num_tokens] positions = buffers.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens]
hidden_states = self.hidden_states[:num_seqs] hidden_states = buffers.hidden_states[:num_seqs]
topk_p = self.topk_p[:num_seqs] topk_p = buffers.topk_p[:num_seqs]
topk_index = self.topk_index[:num_seqs] topk_index = buffers.topk_index[:num_seqs]
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens] * self.dp_size, [num_tokens] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens] * self.dp_size, [num_tokens] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_num_tokens = self.global_num_tokens_gpu global_num_tokens = buffers.global_num_tokens_gpu
global_dp_buffer_len = num_tokens * self.dp_size global_dp_buffer_len = num_tokens * self.dp_size
global_num_tokens_for_logprob = self.global_num_tokens_for_logprob_gpu global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
elif self.require_attn_tp_gather: elif self.require_attn_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens], [num_tokens],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens], [num_tokens],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_num_tokens = self.global_num_tokens_gpu global_num_tokens = buffers.global_num_tokens_gpu
global_dp_buffer_len = num_tokens global_dp_buffer_len = num_tokens
global_num_tokens_for_logprob = self.global_num_tokens_for_logprob_gpu global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
else: else:
global_num_tokens = None global_num_tokens = None
global_dp_buffer_len = None global_dp_buffer_len = None
@@ -319,6 +354,7 @@ class EAGLEDraftCudaGraphRunner:
def replay(self, forward_batch: ForwardBatch): def replay(self, forward_batch: ForwardBatch):
assert forward_batch.out_cache_loc is not None assert forward_batch.out_cache_loc is not None
self.deepep_adapter.replay() self.deepep_adapter.replay()
buffers = self.buffers
raw_bs = forward_batch.batch_size raw_bs = forward_batch.batch_size
raw_num_token = raw_bs * self.num_tokens_per_bs raw_num_token = raw_bs * self.num_tokens_per_bs
@@ -338,40 +374,40 @@ class EAGLEDraftCudaGraphRunner:
bs = self.capture_bs[index] bs = self.capture_bs[index]
if bs != raw_bs: if bs != raw_bs:
self.seq_lens.fill_(self.seq_len_fill_value) buffers.seq_lens.fill_(self.seq_len_fill_value)
self.out_cache_loc.zero_() buffers.out_cache_loc.zero_()
self.positions.zero_() buffers.positions.zero_()
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
# Common inputs # Common inputs
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens) buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
self.out_cache_loc[: raw_num_token * self.speculative_num_steps].copy_( buffers.out_cache_loc[: raw_num_token * self.speculative_num_steps].copy_(
forward_batch.out_cache_loc forward_batch.out_cache_loc
) )
self.positions[:raw_num_token].copy_(forward_batch.positions) buffers.positions[:raw_num_token].copy_(forward_batch.positions)
self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p)
self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index)
self.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
# TODO(ch-wan): support num_token_non_padded # TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer: if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.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) buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
# Attention backend # Attention backend
if bs != raw_bs: if bs != raw_bs:
forward_batch.batch_size = bs forward_batch.batch_size = bs
forward_batch.seq_lens = self.seq_lens[:bs] forward_batch.seq_lens = buffers.seq_lens[:bs]
forward_batch.req_pool_indices = self.req_pool_indices[:bs] forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
forward_batch.positions = self.positions[:num_tokens] forward_batch.positions = buffers.positions[:num_tokens]
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs: if bs != raw_bs:
self.seq_lens_cpu.fill_(self.seq_len_fill_value) buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
forward_batch.seq_lens_cpu = self.seq_lens_cpu[:bs] forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs]
self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph( self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
forward_batch, bs forward_batch, bs
@@ -387,10 +423,10 @@ class EAGLEDraftCudaGraphRunner:
if bs != raw_bs: if bs != raw_bs:
out = self._postprocess_output_to_raw_bs(out, raw_bs) out = self._postprocess_output_to_raw_bs(out, raw_bs)
forward_batch.batch_size = raw_bs forward_batch.batch_size = raw_bs
forward_batch.positions = self.positions[:raw_num_token] forward_batch.positions = buffers.positions[:raw_num_token]
forward_batch.seq_lens = self.seq_lens[:raw_bs] forward_batch.seq_lens = buffers.seq_lens[:raw_bs]
forward_batch.req_pool_indices = self.req_pool_indices[:raw_bs] forward_batch.req_pool_indices = buffers.req_pool_indices[:raw_bs]
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
forward_batch.seq_lens_cpu = self.seq_lens_cpu[:raw_bs] forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
return out return out
@@ -1,7 +1,8 @@
from __future__ import annotations from __future__ import annotations
import bisect import bisect
from typing import TYPE_CHECKING, Callable from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, Optional
import torch import torch
@@ -23,6 +24,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.spec_utils import fast_topk from sglang.srt.speculative.spec_utils import fast_topk
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -36,6 +38,23 @@ if TYPE_CHECKING:
from sglang.srt.speculative.eagle_worker import EAGLEWorker from sglang.srt.speculative.eagle_worker import EAGLEWorker
@dataclass
class EagleDraftExtendInputBuffers(ForwardInputBuffers):
input_ids: torch.Tensor
req_pool_indices: torch.Tensor
out_cache_loc: torch.Tensor
positions: torch.Tensor
mrope_positions: torch.Tensor
hidden_states: torch.Tensor
seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor
extend_seq_lens: torch.Tensor
accept_length: torch.Tensor
next_token_logits_buffer: torch.Tensor
global_num_tokens_gpu: Optional[torch.Tensor]
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
class EAGLEDraftExtendCudaGraphRunner: class EAGLEDraftExtendCudaGraphRunner:
def __init__(self, eagle_worker: EAGLEWorker): def __init__(self, eagle_worker: EAGLEWorker):
# Parse args # Parse args
@@ -80,7 +99,7 @@ class EAGLEDraftExtendCudaGraphRunner:
self.seq_len_fill_value = ( self.seq_len_fill_value = (
self.eagle_worker.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value() self.eagle_worker.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value()
) )
self.seq_lens_cpu = torch.full( seq_lens_cpu = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32 (self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
) )
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
@@ -90,21 +109,19 @@ class EAGLEDraftExtendCudaGraphRunner:
# Graph inputs # Graph inputs
with torch.device(model_runner.device): with torch.device(model_runner.device):
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64) input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32) req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
self.out_cache_loc = torch.ones( out_cache_loc = torch.ones(
(self.max_num_token,), dtype=self._cache_loc_dtype() (self.max_num_token,), dtype=self._cache_loc_dtype()
) )
self.positions = torch.zeros((self.max_num_token,), dtype=torch.int64) positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
self.mrope_positions = torch.zeros( mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
(3, self.max_num_token), dtype=torch.int64
)
if ( if (
self.eagle_worker.speculative_algorithm.is_eagle3() self.eagle_worker.speculative_algorithm.is_eagle3()
and self.eagle_worker.eagle_use_aux_hidden_state and self.eagle_worker.eagle_use_aux_hidden_state
): ):
self.hidden_states = torch.zeros( hidden_states = torch.zeros(
( (
self.max_num_token, self.max_num_token,
( (
@@ -120,40 +137,40 @@ class EAGLEDraftExtendCudaGraphRunner:
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
else: else:
self.hidden_states = torch.zeros( hidden_states = torch.zeros(
(self.max_num_token, self.model_runner.model_config.hidden_size), (self.max_num_token, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
self.seq_len_fill_value = ( self.seq_len_fill_value = (
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value() self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
) )
self.seq_lens = torch.full( seq_lens = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32 (self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
) )
self.extend_seq_lens = torch.full( extend_seq_lens = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
) )
self.accept_length = torch.full( accept_length = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
) )
if self.require_gathered_buffer: if self.require_gathered_buffer:
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu = torch.zeros( global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
else: else:
assert self.require_attn_tp_gather assert self.require_attn_tp_gather
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32) global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(1,), dtype=torch.int32 (1,), dtype=torch.int32
) )
else: else:
self.global_num_tokens_gpu = None global_num_tokens_gpu = None
self.global_num_tokens_for_logprob_gpu = None global_num_tokens_for_logprob_gpu = None
if hasattr( if hasattr(
self.model_runner.model_config.hf_config, "draft_vocab_size" self.model_runner.model_config.hf_config, "draft_vocab_size"
@@ -166,7 +183,7 @@ class EAGLEDraftExtendCudaGraphRunner:
else: else:
vocab_size = self.model_runner.model_config.vocab_size vocab_size = self.model_runner.model_config.vocab_size
self.next_token_logits_buffer = torch.zeros( next_token_logits_buffer = torch.zeros(
( (
( (
self.max_bs * self.num_tokens_per_bs self.max_bs * self.num_tokens_per_bs
@@ -178,6 +195,23 @@ class EAGLEDraftExtendCudaGraphRunner:
dtype=torch.float, dtype=torch.float,
) )
self.buffers = EagleDraftExtendInputBuffers(
input_ids=input_ids,
req_pool_indices=req_pool_indices,
out_cache_loc=out_cache_loc,
positions=positions,
mrope_positions=mrope_positions,
hidden_states=hidden_states,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
extend_seq_lens=extend_seq_lens,
accept_length=accept_length,
next_token_logits_buffer=next_token_logits_buffer,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
)
self.buffers.share_buffers()
# Capture # Capture
try: try:
with model_capture_mode(): with model_capture_mode():
@@ -233,23 +267,24 @@ class EAGLEDraftExtendCudaGraphRunner:
CudaGraphRunner.capture(self) CudaGraphRunner.capture(self)
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0): def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
buffers = self.buffers
graph = self._create_graph() graph = self._create_graph()
stream = self.stream stream = self.stream
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
# Graph inputs # Graph inputs
input_ids = self.input_ids[:num_tokens] input_ids = buffers.input_ids[:num_tokens]
req_pool_indices = self.req_pool_indices[:bs] req_pool_indices = buffers.req_pool_indices[:bs]
seq_lens = self.seq_lens[:bs] seq_lens = buffers.seq_lens[:bs]
seq_lens_cpu = self.seq_lens_cpu[:bs] seq_lens_cpu = buffers.seq_lens_cpu[:bs]
extend_seq_lens = self.extend_seq_lens[:bs] extend_seq_lens = buffers.extend_seq_lens[:bs]
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs] extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs]
out_cache_loc = self.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
positions = self.positions[:num_tokens] positions = buffers.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens]
hidden_states = self.hidden_states[:num_tokens] hidden_states = buffers.hidden_states[:num_tokens]
accept_length = self.accept_length[:bs] accept_length = buffers.accept_length[:bs]
next_token_logits_buffer = self.next_token_logits_buffer[ next_token_logits_buffer = buffers.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens : bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
] ]
@@ -260,34 +295,34 @@ class EAGLEDraftExtendCudaGraphRunner:
) )
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens] * self.dp_size, [num_tokens] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens_for_logprob] * self.dp_size, [num_tokens_for_logprob] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_dp_buffer_len = num_tokens * self.dp_size global_dp_buffer_len = num_tokens * self.dp_size
elif self.require_attn_tp_gather: elif self.require_attn_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens], [num_tokens],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens_for_logprob], [num_tokens_for_logprob],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_dp_buffer_len = num_tokens global_dp_buffer_len = num_tokens
@@ -320,8 +355,8 @@ class EAGLEDraftExtendCudaGraphRunner:
return_logprob=False, return_logprob=False,
positions=positions, positions=positions,
mrope_positions=mrope_positions, mrope_positions=mrope_positions,
global_num_tokens_gpu=self.global_num_tokens_gpu, global_num_tokens_gpu=buffers.global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu, global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(), dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len, global_dp_buffer_len=global_dp_buffer_len,
spec_algorithm=self.model_runner.spec_algorithm, spec_algorithm=self.model_runner.spec_algorithm,
@@ -380,6 +415,7 @@ class EAGLEDraftExtendCudaGraphRunner:
def replay(self, forward_batch: ForwardBatch): def replay(self, forward_batch: ForwardBatch):
assert forward_batch.out_cache_loc is not None assert forward_batch.out_cache_loc is not None
self.deepep_adapter.replay() self.deepep_adapter.replay()
buffers = self.buffers
# batch_size and num_seqs can be different in case there are finished examples # batch_size and num_seqs can be different in case there are finished examples
# in the batch, which will not be counted as num_seqs # in the batch, which will not be counted as num_seqs
@@ -398,45 +434,47 @@ class EAGLEDraftExtendCudaGraphRunner:
bs = self.capture_bs[index] bs = self.capture_bs[index]
if bs * self.num_tokens_per_bs != num_tokens: if bs * self.num_tokens_per_bs != num_tokens:
self.seq_lens.fill_(self.seq_len_fill_value) buffers.seq_lens.fill_(self.seq_len_fill_value)
self.out_cache_loc.zero_() buffers.out_cache_loc.zero_()
self.positions.zero_() buffers.positions.zero_()
self.accept_length.fill_(self.num_tokens_per_bs) buffers.accept_length.fill_(self.num_tokens_per_bs)
self.extend_seq_lens.fill_(self.num_tokens_per_bs) buffers.extend_seq_lens.fill_(self.num_tokens_per_bs)
# Common inputs # Common inputs
self.input_ids[:num_tokens].copy_(forward_batch.input_ids) buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens) buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
if forward_batch.extend_seq_lens is not None: if forward_batch.extend_seq_lens is not None:
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens) buffers.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
else: else:
self.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs) buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc) buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
self.positions[:num_tokens].copy_(forward_batch.positions) buffers.positions[:num_tokens].copy_(forward_batch.positions)
if ( if (
forward_batch.spec_info.hidden_states.shape[1] forward_batch.spec_info.hidden_states.shape[1]
== self.hidden_states.shape[1] == buffers.hidden_states.shape[1]
): ):
self.hidden_states[:num_tokens].copy_(forward_batch.spec_info.hidden_states) buffers.hidden_states[:num_tokens].copy_(
forward_batch.spec_info.hidden_states
)
if forward_batch.spec_info.accept_length is not None: if forward_batch.spec_info.accept_length is not None:
self.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length) buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
# TODO(ch-wan): support num_token_non_padded # TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer: if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
# V1: pruned_states = bs; V2: pruned_states = num_tokens # V1: pruned_states = bs; V2: pruned_states = num_tokens
if self.forward_mode.is_draft_extend_v2(): if self.forward_mode.is_draft_extend_v2():
self.global_num_tokens_for_logprob_gpu.fill_( buffers.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_bs bs * self.num_tokens_per_bs
) )
else: else:
self.global_num_tokens_for_logprob_gpu.fill_(bs) buffers.global_num_tokens_for_logprob_gpu.fill_(bs)
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs: if bs != raw_bs:
self.seq_lens_cpu.fill_(self.seq_len_fill_value) buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
if forward_batch.extend_seq_lens_cpu is not None: if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
@@ -449,22 +487,22 @@ class EAGLEDraftExtendCudaGraphRunner:
forward_batch.spec_info.extend_seq_lens_cpu = list( forward_batch.spec_info.extend_seq_lens_cpu = list(
self.extend_seq_lens_cpu[:bs] self.extend_seq_lens_cpu[:bs]
) )
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs] forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
if bs != raw_bs: if bs != raw_bs:
forward_batch.spec_info.positions = self.positions[:num_tokens] forward_batch.spec_info.positions = buffers.positions[:num_tokens]
forward_batch.spec_info.accept_length = self.accept_length[:bs] forward_batch.spec_info.accept_length = buffers.accept_length[:bs]
self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph( self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
bs=bs, bs=bs,
req_pool_indices=self.req_pool_indices, req_pool_indices=buffers.req_pool_indices,
seq_lens=self.seq_lens, seq_lens=buffers.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value, + (bs - raw_bs) * self.seq_len_fill_value,
encoder_lens=None, encoder_lens=None,
forward_mode=self.forward_mode, forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
seq_lens_cpu=self.seq_lens_cpu, seq_lens_cpu=buffers.seq_lens_cpu,
) )
# Replay # Replay
@@ -477,7 +515,7 @@ class EAGLEDraftExtendCudaGraphRunner:
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not. # DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
unpadding_bs = num_tokens unpadding_bs = num_tokens
elif bs != raw_bs: elif bs != raw_bs:
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs] forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs]
unpadding_bs = raw_bs unpadding_bs = raw_bs
else: else:
unpadding_bs = None unpadding_bs = None
@@ -17,7 +17,8 @@ from __future__ import annotations
import bisect import bisect
import logging import logging
import time import time
from typing import TYPE_CHECKING, Callable from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, List, Optional
import torch import torch
@@ -39,6 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton
from sglang.srt.speculative.spec_utils import fast_topk from sglang.srt.speculative.spec_utils import fast_topk
@@ -59,6 +61,28 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@dataclass
class MultiLayerEagleDraftExtendInputBuffers(ForwardInputBuffers):
# Sliced from shared parent buffers
input_ids: torch.Tensor
out_cache_loc: torch.Tensor
swa_out_cache_loc: torch.Tensor
positions: torch.Tensor
# Shared from parent
seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor
req_pool_indices: torch.Tensor
accept_length: torch.Tensor
# Per-step buffers
extend_seq_lens: torch.Tensor
extend_start_loc: torch.Tensor
mrope_positions: torch.Tensor
hidden_states: torch.Tensor
next_token_logits_buffer: torch.Tensor
global_num_tokens_gpu: Optional[torch.Tensor]
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
class MultiLayerEagleDraftExtendCudaGraphRunner: class MultiLayerEagleDraftExtendCudaGraphRunner:
def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int): def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int):
# Parse args # Parse args
@@ -109,7 +133,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
next_cuda_graph_runner, next_cuda_graph_runner,
): ):
self.next_cuda_graph_runner = next_cuda_graph_runner self.next_cuda_graph_runner = next_cuda_graph_runner
self.seq_lens_cpu = cuda_graph_buffers["seq_lens_cpu"] seq_lens_cpu = cuda_graph_buffers["seq_lens_cpu"]
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
if self.enable_torch_compile: if self.enable_torch_compile:
@@ -119,62 +143,60 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
with torch.device(self.model_runner.device): with torch.device(self.model_runner.device):
# sliced buffers # sliced buffers
# slice according to max_num_token # slice according to max_num_token
self.input_ids = cuda_graph_buffers["input_ids"][ input_ids = cuda_graph_buffers["input_ids"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
self.out_cache_loc = cuda_graph_buffers["out_cache_loc"][ out_cache_loc = cuda_graph_buffers["out_cache_loc"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
self.swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][ swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
self.positions = cuda_graph_buffers["positions"][ positions = cuda_graph_buffers["positions"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
# shared states # shared states
self.seq_lens = cuda_graph_buffers["seq_lens"] seq_lens = cuda_graph_buffers["seq_lens"]
self.req_pool_indices = cuda_graph_buffers["req_pool_indices"] req_pool_indices = cuda_graph_buffers["req_pool_indices"]
self.accept_length = cuda_graph_buffers["accept_length"] accept_length = cuda_graph_buffers["accept_length"]
self.extend_seq_lens = torch.full( extend_seq_lens = torch.full(
(self.max_bs,), (self.max_bs,),
self.num_tokens_per_bs, self.num_tokens_per_bs,
dtype=torch.int32, dtype=torch.int32,
) )
self.extend_start_loc = torch.arange( extend_start_loc = torch.arange(
0, 0,
self.max_bs * self.num_tokens_per_bs, self.max_bs * self.num_tokens_per_bs,
step=self.num_tokens_per_bs, step=self.num_tokens_per_bs,
dtype=torch.int32, dtype=torch.int32,
) )
self.mrope_positions = torch.zeros( mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
(3, self.max_num_token), dtype=torch.int64
)
self.hidden_states = torch.zeros( hidden_states = torch.zeros(
(self.max_num_token, self.model_runner.model_config.hidden_size), (self.max_num_token, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype, dtype=self.model_runner.dtype,
) )
if self.require_gathered_buffer: if self.require_gathered_buffer:
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu = torch.zeros( global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32 (self.dp_size,), dtype=torch.int32
) )
else: else:
assert self.require_attn_tp_gather assert self.require_attn_tp_gather
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32) global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
self.global_num_tokens_for_logprob_gpu = torch.zeros( global_num_tokens_for_logprob_gpu = torch.zeros(
(1,), dtype=torch.int32 (1,), dtype=torch.int32
) )
else: else:
self.global_num_tokens_gpu = None global_num_tokens_gpu = None
self.global_num_tokens_for_logprob_gpu = None global_num_tokens_for_logprob_gpu = None
if hasattr( if hasattr(
self.model_runner.model_config.hf_config, "draft_vocab_size" self.model_runner.model_config.hf_config, "draft_vocab_size"
@@ -187,7 +209,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
else: else:
vocab_size = self.model_runner.model_config.vocab_size vocab_size = self.model_runner.model_config.vocab_size
self.next_token_logits_buffer = torch.zeros( next_token_logits_buffer = torch.zeros(
( (
( (
self.max_bs * self.num_tokens_per_bs self.max_bs * self.num_tokens_per_bs
@@ -199,6 +221,24 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
dtype=torch.float, dtype=torch.float,
) )
self.buffers = MultiLayerEagleDraftExtendInputBuffers(
input_ids=input_ids,
out_cache_loc=out_cache_loc,
swa_out_cache_loc=swa_out_cache_loc,
positions=positions,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
req_pool_indices=req_pool_indices,
accept_length=accept_length,
extend_seq_lens=extend_seq_lens,
extend_start_loc=extend_start_loc,
mrope_positions=mrope_positions,
hidden_states=hidden_states,
next_token_logits_buffer=next_token_logits_buffer,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
)
# Capture # Capture
try: try:
with model_capture_mode(): with model_capture_mode():
@@ -250,54 +290,55 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
CudaGraphRunner.capture(self) CudaGraphRunner.capture(self)
def get_forward_batch(self, bs: int) -> ForwardBatch: def get_forward_batch(self, bs: int) -> ForwardBatch:
buffers = self.buffers
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
# Graph inputs # Graph inputs
input_ids = self.input_ids[:num_tokens] input_ids = buffers.input_ids[:num_tokens]
req_pool_indices = self.req_pool_indices[:bs] req_pool_indices = buffers.req_pool_indices[:bs]
seq_lens = self.seq_lens[:bs] seq_lens = buffers.seq_lens[:bs]
seq_lens_cpu = self.seq_lens_cpu[:bs] seq_lens_cpu = buffers.seq_lens_cpu[:bs]
extend_seq_lens = self.extend_seq_lens[:bs] extend_seq_lens = buffers.extend_seq_lens[:bs]
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs] extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs]
extend_start_loc = self.extend_start_loc[:bs] extend_start_loc = buffers.extend_start_loc[:bs]
accept_length = self.accept_length[:bs] accept_length = buffers.accept_length[:bs]
out_cache_loc = self.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
positions = self.positions[:num_tokens] positions = buffers.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens]
hidden_states = self.hidden_states[:num_tokens] hidden_states = buffers.hidden_states[:num_tokens]
next_token_logits_buffer = self.next_token_logits_buffer[ next_token_logits_buffer = buffers.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens : bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
] ]
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens] * self.dp_size, [num_tokens] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens] * self.dp_size, [num_tokens] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_dp_buffer_len = num_tokens * self.dp_size global_dp_buffer_len = num_tokens * self.dp_size
elif self.require_attn_tp_gather: elif self.require_attn_tp_gather:
self.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
[num_tokens], [num_tokens],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
self.global_num_tokens_for_logprob_gpu.copy_( buffers.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[bs], [bs],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=buffers.input_ids.device,
) )
) )
global_dp_buffer_len = num_tokens global_dp_buffer_len = num_tokens
@@ -326,8 +367,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
return_logprob=False, return_logprob=False,
positions=positions, positions=positions,
mrope_positions=mrope_positions, mrope_positions=mrope_positions,
global_num_tokens_gpu=self.global_num_tokens_gpu, global_num_tokens_gpu=buffers.global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu, global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(), dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len, global_dp_buffer_len=global_dp_buffer_len,
spec_algorithm=self.model_runner.spec_algorithm, spec_algorithm=self.model_runner.spec_algorithm,
@@ -346,6 +387,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
return forward_batch return forward_batch
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0): def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
buffers = self.buffers
graph = self._create_graph() graph = self._create_graph()
stream = self.stream stream = self.stream
@@ -390,7 +432,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
select_index = ( select_index = (
torch.arange(bs, device=self.model_runner.device) torch.arange(bs, device=self.model_runner.device)
* (self.speculative_num_draft_tokens + self.step) * (self.speculative_num_draft_tokens + self.step)
+ self.accept_length[:bs] + buffers.accept_length[:bs]
- 1 - 1
+ self.step + self.step
) )
@@ -399,24 +441,25 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
ret.topk_p, ret.topk_index = fast_topk(probs, self.topk, dim=-1) ret.topk_p, ret.topk_index = fast_topk(probs, self.topk, dim=-1)
if self.next_cuda_graph_runner is not None: if self.next_cuda_graph_runner is not None:
next_buffers = self.next_cuda_graph_runner.buffers
padding_lens = ( padding_lens = (
self.speculative_num_draft_tokens - self.accept_length[:bs] self.speculative_num_draft_tokens - buffers.accept_length[:bs]
) )
assign_new_state_triton( assign_new_state_triton(
ret.topk_index, ret.topk_index,
self.input_ids, buffers.input_ids,
self.positions, buffers.positions,
self.hidden_states, buffers.hidden_states,
self.out_cache_loc, buffers.out_cache_loc,
self.extend_seq_lens, buffers.extend_seq_lens,
self.extend_start_loc, buffers.extend_start_loc,
self.next_cuda_graph_runner.input_ids, next_buffers.input_ids,
self.next_cuda_graph_runner.positions, next_buffers.positions,
self.next_cuda_graph_runner.hidden_states, next_buffers.hidden_states,
self.next_cuda_graph_runner.out_cache_loc, next_buffers.out_cache_loc,
self.next_cuda_graph_runner.extend_seq_lens, next_buffers.extend_seq_lens,
self.next_cuda_graph_runner.extend_start_loc, next_buffers.extend_start_loc,
self.next_cuda_graph_runner.seq_lens, next_buffers.seq_lens,
padding_lens, padding_lens,
forward_batch.batch_size, forward_batch.batch_size,
self.step, self.step,
@@ -424,9 +467,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
forward_batch.req_to_token_pool.req_to_token, forward_batch.req_to_token_pool.req_to_token,
self.eagle_worker.req_to_hidden_states_pool, self.eagle_worker.req_to_hidden_states_pool,
) )
self.next_cuda_graph_runner.swa_out_cache_loc.copy_( next_buffers.swa_out_cache_loc.copy_(
self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa( self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa(
self.next_cuda_graph_runner.out_cache_loc next_buffers.out_cache_loc
) )
) )
@@ -446,27 +489,30 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
def init_replay_state( def init_replay_state(
self, forward_batch: ForwardBatch, bs: int, raw_bs: int, num_tokens: int self, forward_batch: ForwardBatch, bs: int, raw_bs: int, num_tokens: int
): ):
buffers = self.buffers
# Common inputs # Common inputs
self.input_ids[:num_tokens].copy_(forward_batch.input_ids) buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens) buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
if forward_batch.extend_seq_lens is not None: if forward_batch.extend_seq_lens is not None:
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens) buffers.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
self.extend_start_loc[:raw_bs].copy_(forward_batch.extend_start_loc) buffers.extend_start_loc[:raw_bs].copy_(forward_batch.extend_start_loc)
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc) buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
self.positions[:num_tokens].copy_(forward_batch.positions) buffers.positions[:num_tokens].copy_(forward_batch.positions)
if ( if (
forward_batch.spec_info.hidden_states.shape[1] forward_batch.spec_info.hidden_states.shape[1]
== self.hidden_states.shape[1] == buffers.hidden_states.shape[1]
): ):
self.hidden_states[:num_tokens].copy_(forward_batch.spec_info.hidden_states) buffers.hidden_states[:num_tokens].copy_(
forward_batch.spec_info.hidden_states
)
if forward_batch.spec_info.accept_length is not None: if forward_batch.spec_info.accept_length is not None:
self.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length) buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs: if bs != raw_bs:
self.seq_lens_cpu.fill_(self.seq_len_fill_value) buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
if forward_batch.extend_seq_lens_cpu is not None: if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
@@ -474,6 +520,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
def replay(self, forward_batch: ForwardBatch, init_state: bool = True): def replay(self, forward_batch: ForwardBatch, init_state: bool = True):
assert forward_batch.out_cache_loc is not None assert forward_batch.out_cache_loc is not None
self.deepep_adapter.replay() self.deepep_adapter.replay()
buffers = self.buffers
# batch_size and num_seqs can be different in case there are finished examples # batch_size and num_seqs can be different in case there are finished examples
# in the batch, which will not be counted as num_seqs # in the batch, which will not be counted as num_seqs
@@ -492,28 +539,28 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
self.init_replay_state(forward_batch, bs, raw_bs, num_tokens) self.init_replay_state(forward_batch, bs, raw_bs, num_tokens)
if self.require_gathered_buffer: if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) buffers.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) buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
forward_batch.spec_info.hidden_states = self.hidden_states[:num_tokens] forward_batch.spec_info.hidden_states = buffers.hidden_states[:num_tokens]
forward_batch.spec_info.accept_length = self.accept_length[:bs] forward_batch.spec_info.accept_length = buffers.accept_length[:bs]
forward_batch.spec_info.num_tokens_per_req = self.num_tokens_per_bs forward_batch.spec_info.num_tokens_per_req = self.num_tokens_per_bs
forward_batch.spec_info.num_tokens_for_logprob_per_req = 1 forward_batch.spec_info.num_tokens_for_logprob_per_req = 1
forward_batch.spec_info.positions = self.positions[:num_tokens] forward_batch.spec_info.positions = buffers.positions[:num_tokens]
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs] forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
self.eagle_worker.draft_extend_attn_backend_list[ self.eagle_worker.draft_extend_attn_backend_list[
self.step self.step
].init_forward_metadata_replay_cuda_graph( ].init_forward_metadata_replay_cuda_graph(
bs=bs, bs=bs,
req_pool_indices=self.req_pool_indices, req_pool_indices=buffers.req_pool_indices,
seq_lens=self.seq_lens, seq_lens=buffers.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value, + (bs - raw_bs) * self.seq_len_fill_value,
encoder_lens=None, encoder_lens=None,
forward_mode=self.forward_mode, forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
seq_lens_cpu=self.seq_lens_cpu, seq_lens_cpu=buffers.seq_lens_cpu,
) )
# Replay # Replay
@@ -526,7 +573,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not. # DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
unpadding_bs = num_tokens unpadding_bs = num_tokens
elif bs != raw_bs: elif bs != raw_bs:
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs] forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs]
unpadding_bs = raw_bs unpadding_bs = raw_bs
else: else:
unpadding_bs = None unpadding_bs = None
@@ -565,8 +612,8 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.runners = [None] * self.speculative_num_steps self.runners = [None] * self.speculative_num_steps
return return
self.runners = [] self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = []
buffer_len_list = [] buffer_len_list: List[int] = []
# 1. Capture loop # 1. Capture loop
for step in range(self.speculative_num_steps): for step in range(self.speculative_num_steps):
@@ -498,13 +498,13 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
self.cuda_graph_runner_for_draft_extend.get_last_runner() self.cuda_graph_runner_for_draft_extend.get_last_runner()
) )
assign_hidden_states_pool_triton( assign_hidden_states_pool_triton(
last_cuda_graph_runner.hidden_states, last_cuda_graph_runner.buffers.hidden_states,
last_cuda_graph_runner.req_pool_indices, last_cuda_graph_runner.buffers.req_pool_indices,
self.req_to_hidden_states_pool, self.req_to_hidden_states_pool,
self.speculative_num_steps - 1, self.speculative_num_steps - 1,
forward_batch.batch_size, forward_batch.batch_size,
last_cuda_graph_runner.extend_seq_lens, last_cuda_graph_runner.buffers.extend_seq_lens,
last_cuda_graph_runner.extend_start_loc, last_cuda_graph_runner.buffers.extend_start_loc,
) )
# Reorganize the spec info for the next batch # Reorganize the spec info for the next batch