Pre-set SWA cache location in CudaGraphRunner (#23552)
This commit is contained in:
@@ -138,6 +138,7 @@ 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
|
||||||
@@ -171,6 +172,7 @@ 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,
|
||||||
) -> "DecodeInputBuffers":
|
) -> "DecodeInputBuffers":
|
||||||
with torch.device(device):
|
with torch.device(device):
|
||||||
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
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)
|
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
||||||
seq_lens = torch.full((max_bs,), seq_len_fill_value, 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)
|
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)
|
||||||
@@ -249,6 +256,7 @@ 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,
|
||||||
@@ -356,6 +364,14 @@ 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)
|
||||||
|
|
||||||
@@ -685,6 +701,7 @@ class CudaGraphRunner:
|
|||||||
ne_token_table=(
|
ne_token_table=(
|
||||||
model_runner.token_table if self.use_ngram_embedding else None
|
model_runner.token_table if self.use_ngram_embedding else None
|
||||||
),
|
),
|
||||||
|
is_hybrid_swa=model_runner.is_hybrid_swa,
|
||||||
)
|
)
|
||||||
self.buffers.share_buffers()
|
self.buffers.share_buffers()
|
||||||
|
|
||||||
@@ -1131,6 +1148,15 @@ class CudaGraphRunner:
|
|||||||
|
|
||||||
self.deepep_adapter.capture(is_extend_in_batch=False)
|
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):
|
for _ in range(2):
|
||||||
self.device_module.synchronize()
|
self.device_module.synchronize()
|
||||||
self.model_runner.tp_group.barrier()
|
self.model_runner.tp_group.barrier()
|
||||||
@@ -1218,6 +1244,7 @@ class CudaGraphRunner:
|
|||||||
),
|
),
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
pp_proxy_tensors=pp_proxy_tensors,
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.model_runner.spec_algorithm.is_dflash()
|
self.model_runner.spec_algorithm.is_dflash()
|
||||||
and self.model_runner.is_draft_worker
|
and self.model_runner.is_draft_worker
|
||||||
|
|||||||
Reference in New Issue
Block a user