Pre-set SWA cache location in CudaGraphRunner (#23552)

This commit is contained in:
Lianmin Zheng
2026-04-23 16:51:29 -07:00
committed by GitHub
parent bb962b0046
commit 95d021b523
@@ -138,6 +138,7 @@ 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
@@ -171,6 +172,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
is_hybrid_swa: bool = False,
) -> "DecodeInputBuffers":
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
@@ -178,6 +180,11 @@ 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.int32)
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)
@@ -249,6 +256,7 @@ 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,
@@ -356,6 +364,14 @@ 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)
@@ -685,6 +701,7 @@ class CudaGraphRunner:
ne_token_table=(
model_runner.token_table if self.use_ngram_embedding else None
),
is_hybrid_swa=model_runner.is_hybrid_swa,
)
self.buffers.share_buffers()
@@ -1131,6 +1148,15 @@ class CudaGraphRunner:
self.deepep_adapter.capture(is_extend_in_batch=False)
# swa_loc must be set before capture so that set_kv_buffer's
# Python branch (if self.swa_loc is not None) takes the fast path,
# and the graph records GPU ops using this buffer instead of the
# per-layer translate_loc_from_full_to_swa fallback.
if self.buffers.out_cache_loc_swa is not None:
self.model_runner.token_to_kv_pool.set_swa_loc(
self.buffers.out_cache_loc_swa[:num_tokens]
)
for _ in range(2):
self.device_module.synchronize()
self.model_runner.tp_group.barrier()
@@ -1218,6 +1244,7 @@ class CudaGraphRunner:
),
pp_proxy_tensors=pp_proxy_tensors,
)
if (
self.model_runner.spec_algorithm.is_dflash()
and self.model_runner.is_draft_worker