[Refactor] Remove dead out_cache_loc_swa buffers (#28968)
This commit is contained in:
@@ -172,7 +172,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
max_bs=self.max_bs,
|
max_bs=self.max_bs,
|
||||||
max_num_tokens=self.max_num_tokens,
|
max_num_tokens=self.max_num_tokens,
|
||||||
cache_loc_dtype=self._cache_loc_dtype(),
|
cache_loc_dtype=self._cache_loc_dtype(),
|
||||||
is_hybrid_swa=model_runner.is_hybrid_swa,
|
|
||||||
is_multimodal=self.is_multimodal,
|
is_multimodal=self.is_multimodal,
|
||||||
hidden_size=self.model_runner.model_config.hidden_size,
|
hidden_size=self.model_runner.model_config.hidden_size,
|
||||||
dtype=self.model_runner.dtype,
|
dtype=self.model_runner.dtype,
|
||||||
|
|||||||
@@ -69,7 +69,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
seq_lens: torch.Tensor
|
seq_lens: torch.Tensor
|
||||||
seq_lens_cpu: torch.Tensor
|
seq_lens_cpu: torch.Tensor
|
||||||
out_cache_loc: torch.Tensor
|
out_cache_loc: torch.Tensor
|
||||||
out_cache_loc_swa: Optional[torch.Tensor]
|
|
||||||
positions: torch.Tensor
|
positions: torch.Tensor
|
||||||
mrope_positions: torch.Tensor
|
mrope_positions: torch.Tensor
|
||||||
num_token_non_padded: torch.Tensor
|
num_token_non_padded: torch.Tensor
|
||||||
@@ -105,7 +104,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
cache_loc_dtype: torch.dtype,
|
cache_loc_dtype: torch.dtype,
|
||||||
enable_mamba_track: bool,
|
enable_mamba_track: bool,
|
||||||
ne_token_table: Optional[torch.Tensor] = None,
|
ne_token_table: Optional[torch.Tensor] = None,
|
||||||
is_hybrid_swa: bool = False,
|
|
||||||
hc_hidden_size: Optional[int] = None,
|
hc_hidden_size: Optional[int] = None,
|
||||||
pp_proxy_topk_size: Optional[int] = None,
|
pp_proxy_topk_size: Optional[int] = None,
|
||||||
) -> DecodeInputBuffers:
|
) -> DecodeInputBuffers:
|
||||||
@@ -115,11 +113,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
||||||
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
|
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
|
||||||
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
||||||
out_cache_loc_swa = (
|
|
||||||
torch.zeros((max_num_token,), dtype=torch.int64)
|
|
||||||
if is_hybrid_swa
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
positions = 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)
|
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
|
||||||
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
|
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
|
||||||
@@ -206,7 +199,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
out_cache_loc_swa=out_cache_loc_swa,
|
|
||||||
positions=positions,
|
positions=positions,
|
||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
num_token_non_padded=num_token_non_padded,
|
num_token_non_padded=num_token_non_padded,
|
||||||
@@ -326,14 +318,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
dsts.append(buf[:dim])
|
dsts.append(buf[:dim])
|
||||||
srcs.append(src)
|
srcs.append(src)
|
||||||
|
|
||||||
# SWA cache location (int32, separate from the int64 batch above).
|
|
||||||
if (
|
|
||||||
self.out_cache_loc_swa is not None
|
|
||||||
and forward_batch.out_cache_loc_swa is not None
|
|
||||||
):
|
|
||||||
dsts.append(self.out_cache_loc_swa[:raw_num_token])
|
|
||||||
srcs.append(forward_batch.out_cache_loc_swa[:raw_num_token])
|
|
||||||
|
|
||||||
# Batch all GPU copies, grouped by dtype pair.
|
# Batch all GPU copies, grouped by dtype pair.
|
||||||
_grouped_foreach_copy_(dsts, srcs)
|
_grouped_foreach_copy_(dsts, srcs)
|
||||||
|
|
||||||
@@ -347,7 +331,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
class PrefillInputBuffers(ForwardInputBuffers):
|
class PrefillInputBuffers(ForwardInputBuffers):
|
||||||
input_ids: torch.Tensor
|
input_ids: torch.Tensor
|
||||||
out_cache_loc: torch.Tensor
|
out_cache_loc: torch.Tensor
|
||||||
out_cache_loc_swa: Optional[torch.Tensor]
|
|
||||||
mamba_track_indices: Optional[torch.Tensor]
|
mamba_track_indices: Optional[torch.Tensor]
|
||||||
mamba_track_mask: Optional[torch.Tensor]
|
mamba_track_mask: Optional[torch.Tensor]
|
||||||
mamba_track_seqlens: Optional[torch.Tensor]
|
mamba_track_seqlens: Optional[torch.Tensor]
|
||||||
@@ -363,7 +346,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
|
|||||||
max_bs: int,
|
max_bs: int,
|
||||||
max_num_tokens: int,
|
max_num_tokens: int,
|
||||||
cache_loc_dtype: torch.dtype,
|
cache_loc_dtype: torch.dtype,
|
||||||
is_hybrid_swa: bool,
|
|
||||||
is_multimodal: bool,
|
is_multimodal: bool,
|
||||||
hidden_size: int,
|
hidden_size: int,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
@@ -372,11 +354,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
|
|||||||
with torch.device(device):
|
with torch.device(device):
|
||||||
input_ids = torch.zeros((max_num_tokens,), dtype=torch.int64)
|
input_ids = torch.zeros((max_num_tokens,), dtype=torch.int64)
|
||||||
out_cache_loc = torch.zeros((max_num_tokens,), dtype=cache_loc_dtype)
|
out_cache_loc = torch.zeros((max_num_tokens,), dtype=cache_loc_dtype)
|
||||||
out_cache_loc_swa = (
|
|
||||||
torch.zeros((max_num_tokens,), dtype=torch.int64)
|
|
||||||
if is_hybrid_swa
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
mamba_track_indices = (
|
mamba_track_indices = (
|
||||||
torch.zeros((max_bs,), dtype=torch.int64)
|
torch.zeros((max_bs,), dtype=torch.int64)
|
||||||
if enable_mamba_track
|
if enable_mamba_track
|
||||||
@@ -402,7 +379,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
|
|||||||
return cls(
|
return cls(
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
out_cache_loc_swa=out_cache_loc_swa,
|
|
||||||
mamba_track_indices=mamba_track_indices,
|
mamba_track_indices=mamba_track_indices,
|
||||||
mamba_track_mask=mamba_track_mask,
|
mamba_track_mask=mamba_track_mask,
|
||||||
mamba_track_seqlens=mamba_track_seqlens,
|
mamba_track_seqlens=mamba_track_seqlens,
|
||||||
|
|||||||
Reference in New Issue
Block a user