[Refactor] Remove dead out_cache_loc_swa buffers (#28968)

This commit is contained in:
Cheng Wan
2026-06-23 00:03:30 -07:00
committed by GitHub
parent 52a90c9a36
commit c4376aaa88
2 changed files with 0 additions and 25 deletions
@@ -172,7 +172,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
max_bs=self.max_bs,
max_num_tokens=self.max_num_tokens,
cache_loc_dtype=self._cache_loc_dtype(),
is_hybrid_swa=model_runner.is_hybrid_swa,
is_multimodal=self.is_multimodal,
hidden_size=self.model_runner.model_config.hidden_size,
dtype=self.model_runner.dtype,
@@ -69,7 +69,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor
out_cache_loc: torch.Tensor
out_cache_loc_swa: Optional[torch.Tensor]
positions: torch.Tensor
mrope_positions: torch.Tensor
num_token_non_padded: torch.Tensor
@@ -105,7 +104,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
is_hybrid_swa: bool = False,
hc_hidden_size: Optional[int] = None,
pp_proxy_topk_size: Optional[int] = None,
) -> DecodeInputBuffers:
@@ -115,11 +113,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
req_pool_indices = torch.zeros((max_bs,), 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_swa = (
torch.zeros((max_num_token,), dtype=torch.int64)
if is_hybrid_swa
else None
)
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)
@@ -206,7 +199,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
positions=positions,
mrope_positions=mrope_positions,
num_token_non_padded=num_token_non_padded,
@@ -326,14 +318,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
dsts.append(buf[:dim])
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.
_grouped_foreach_copy_(dsts, srcs)
@@ -347,7 +331,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
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]
@@ -363,7 +346,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
max_bs: int,
max_num_tokens: int,
cache_loc_dtype: torch.dtype,
is_hybrid_swa: bool,
is_multimodal: bool,
hidden_size: int,
dtype: torch.dtype,
@@ -372,11 +354,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
with torch.device(device):
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_swa = (
torch.zeros((max_num_tokens,), dtype=torch.int64)
if is_hybrid_swa
else None
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64)
if enable_mamba_track
@@ -402,7 +379,6 @@ class PrefillInputBuffers(ForwardInputBuffers):
return cls(
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,