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