fix(swa): eliminate spurious translate_loc_from_full_to_swa warning in BCG and CG paths (#26152)
Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
@@ -329,7 +329,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
seq_lens,
|
||||
req_pool_indices,
|
||||
bs,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
token_to_kv_pool=self.token_to_kv_pool,
|
||||
window_kv_indices=self.cuda_graph_window_kv_indices,
|
||||
)
|
||||
return kv_indptr, window_kv_indptr, window_kv_lens
|
||||
@@ -374,7 +374,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
seq_lens[:bs],
|
||||
req_pool_indices,
|
||||
bs,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
token_to_kv_pool=self.token_to_kv_pool,
|
||||
window_kv_indices=window_kv_indices,
|
||||
)
|
||||
)
|
||||
@@ -461,7 +461,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
forward_batch.req_pool_indices,
|
||||
bs,
|
||||
self.device,
|
||||
self.token_to_kv_pool_allocator,
|
||||
self.token_to_kv_pool,
|
||||
)
|
||||
)
|
||||
window_num_kv_splits = torch.empty(
|
||||
@@ -532,7 +532,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
forward_batch.req_pool_indices,
|
||||
bs,
|
||||
self.device,
|
||||
self.token_to_kv_pool_allocator,
|
||||
self.token_to_kv_pool,
|
||||
)
|
||||
|
||||
custom_mask = spec_info.custom_mask
|
||||
@@ -592,7 +592,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
forward_batch.req_pool_indices,
|
||||
bs,
|
||||
self.device,
|
||||
self.token_to_kv_pool_allocator,
|
||||
self.token_to_kv_pool,
|
||||
)
|
||||
|
||||
qo_indptr = self.qo_indptr
|
||||
@@ -1445,7 +1445,7 @@ def update_sliding_window_buffer(
|
||||
req_pool_indices,
|
||||
bs,
|
||||
device=None,
|
||||
token_to_kv_pool_allocator=None,
|
||||
token_to_kv_pool=None,
|
||||
window_kv_indices=None,
|
||||
):
|
||||
"""Fill window KV buffers for sliding-window attention.
|
||||
@@ -1474,11 +1474,14 @@ def update_sliding_window_buffer(
|
||||
window_kv_indices,
|
||||
req_to_token.stride(0),
|
||||
)
|
||||
if hasattr(token_to_kv_pool_allocator, "translate_loc_from_full_to_swa"):
|
||||
if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"):
|
||||
kv_last_index = window_kv_indptr[-1]
|
||||
# Flush before+after: window_kv_indices is a different tensor than out_cache_loc.
|
||||
token_to_kv_pool.invalidate_loc_cache()
|
||||
window_kv_indices[:kv_last_index] = (
|
||||
token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
|
||||
token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||
window_kv_indices[:kv_last_index]
|
||||
)
|
||||
)
|
||||
token_to_kv_pool.invalidate_loc_cache()
|
||||
return window_kv_indptr, window_kv_indices, window_kv_lens, window_kv_start_idx
|
||||
|
||||
@@ -3266,6 +3266,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.hisparse_coordinator.wait_for_pending_backup()
|
||||
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
|
||||
|
||||
if self.is_hybrid_swa:
|
||||
self.token_to_kv_pool.invalidate_loc_cache()
|
||||
|
||||
# Replay cuda graph if applicable
|
||||
if can_run_graph:
|
||||
ret = self.graph_runner.replay(
|
||||
@@ -3292,9 +3295,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
server_args=self.server_args,
|
||||
)
|
||||
|
||||
if self.is_hybrid_swa:
|
||||
self.token_to_kv_pool.invalidate_loc_cache()
|
||||
|
||||
# Hisparse coordinator — backends now read it from self.model_runner.
|
||||
if self.hisparse_coordinator is not None:
|
||||
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
|
||||
|
||||
Reference in New Issue
Block a user