Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 13 additions & 6 deletions tests/ut/distributed/ascend_store/test_pool_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -899,9 +899,14 @@ def create_recv_thread(*args, **kwargs):
kwargs["invalid_block_ids"].add(7)
self.assertEqual(worker.get_block_ids_with_load_errors(), {7})

def test_wait_for_save_waits_for_save(self):
def test_wait_for_save_submits_batch_without_joining_queue(self):
worker = self._make_worker()
worker.kv_send_thread = MagicMock()
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import KVCacheStoreSendingThread

worker.kv_send_thread = MagicMock(spec=KVCacheStoreSendingThread)
worker.kv_send_thread.request_queue = MagicMock()
save_batch = MagicMock()
worker.kv_send_thread.add_save_batch.return_value = save_batch

req = ReqMeta(
req_id="r1",
Expand All @@ -912,10 +917,12 @@ def test_wait_for_save_waits_for_save(self):
)
meta = AscendConnectorMetadata(set(), set())
meta.add_request(req)
worker.wait_for_save(meta)
worker.kv_send_thread.add_stored_request.assert_called_with("r1")
worker.kv_send_thread.add_request.assert_called_once()
worker.kv_send_thread.request_queue.join.assert_called_once()
module = "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_worker"
with patch(f"{module}.torch.npu", create=True):
worker.wait_for_save(meta)
worker.kv_send_thread.add_save_batch.assert_called_once_with([req])
worker.kv_send_thread.request_queue.join.assert_not_called()
self.assertIs(worker._previous_save_batch, save_batch)

def test_wait_for_save_skip_non_save(self):
worker = self._make_worker()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -241,6 +241,15 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):
assert self.connector_worker is not None
self.connector_worker.register_kv_caches(kv_caches)

def handle_preemptions(self, kv_connector_metadata: KVConnectorMetadata) -> None:
"""Fence the previous save before this step can reuse KV blocks.

This hook is temporarily reused for deferred KV cache save
synchronization and will be replaced by a dedicated mechanism.
"""
assert self.connector_worker is not None
self.connector_worker.wait_for_previous_save()

def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
assert self.connector_worker is not None
self._mamba_copy_bufs = None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -859,6 +859,13 @@ def _decode_adaptor_prefill_pp(
return self.token_database.decode_adaptor_prefill_pp(keys, addrs, sizes)


class KVCacheStoreBatch:
"""FIFO fence for a batch of asynchronous KV store requests."""

def __init__(self) -> None:
self.done = threading.Event()


class KVCacheStoreSendingThread(KVTransferThread):
def __init__(
self,
Expand Down Expand Up @@ -890,6 +897,26 @@ def __init__(
self.completed_events: dict[int, int] = {}
self.worker = worker

def add_stored_request(self, req_id: str):
with self.done_task_lock:
# A later chunk of the same request starts a new save lifecycle.
# Do not let a completion from an earlier chunk release it early.
self.finished_requests.discard(req_id)
self.stored_requests[req_id] += 1

def add_save_batch(self, requests: list[ReqMeta]) -> KVCacheStoreBatch:
"""Queue requests followed by a fence that completes after the batch."""
save_batch = KVCacheStoreBatch()
# Register the entire batch before exposing any request to the send
# thread. Otherwise duplicate req_ids could transiently reach zero and
# be reported as finished between two chunks in the same batch.
for request in requests:
self.add_stored_request(request.req_id)
for request in requests:
self.request_queue.put(request)
self.request_queue.put(save_batch)
return save_batch

def is_stored_request(self, req_id: str) -> bool:
with self.done_task_lock:
return req_id in self.stored_requests
Expand Down Expand Up @@ -919,7 +946,12 @@ def _handle_request_exception(self, request_data: Any):
self.dec_stored_request(req_id)
self.request_queue.task_done()

def _handle_request(self, req_meta: ReqMeta):
def _handle_request(self, req_meta: ReqMeta | KVCacheStoreBatch):
if isinstance(req_meta, KVCacheStoreBatch):
req_meta.done.set()
self.request_queue.task_done()
return

if self.worker is not None and getattr(self.worker, "tp_mismatch", False):
req_id = req_meta.req_id
try:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
)
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator import AscendStoreCoordinator
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import (
KVCacheStoreBatch,
KVCacheStoreKeyLayerRecvingThread,
KVCacheStoreKeyLayerSendingThread,
KVCacheStoreLayerRecvingThread,
Expand Down Expand Up @@ -105,6 +106,7 @@
MEMCACHE_UNMATCHED_STATE = -3101
PARTIAL_LEASE_RETRY_COUNT = 10
PARTIAL_LEASE_RETRY_INTERVAL_S = 0.001
SAVE_BATCH_FAILURE_POLL_INTERVAL_S = 1.0


class KVPoolWorker:
Expand Down Expand Up @@ -396,6 +398,7 @@ def _init_kv_events(self, vllm_config) -> None:
def _init_state_vars(self) -> None:
self.kv_send_thread: KVTransferThread | None = None
self.kv_recv_thread: KVTransferThread | None = None
self._previous_save_batch: KVCacheStoreBatch | None = None
self._transfer_threads_started = False
self.external_slot_release_waiter: Callable[[int], None] | None = None
# Per-rank GVA cache: maps per-rank store key to its allocated GVA.
Expand Down Expand Up @@ -2400,10 +2403,31 @@ def save_kv_layer(self, connector_metadata: AscendConnectorMetadata) -> None:

self.current_layer = self.current_layer + 1

def wait_for_save(self, connector_metadata: AscendConnectorMetadata):
def wait_for_previous_save(self) -> None:
save_batch = self._previous_save_batch
if save_batch is None:
return

assert self.kv_send_thread is not None
send_thread = self.kv_send_thread
wait_start = time.perf_counter()
while True:
send_thread.raise_if_failed()
Comment thread
LCAIZJ marked this conversation as resolved.
if save_batch.done.wait(timeout=SAVE_BATCH_FAILURE_POLL_INTERVAL_S):
break
elapsed = time.perf_counter() - wait_start
logger.debug(
"Previous KV save batch completed after waiting %.3f ms tp_rank=%d",
elapsed * 1000,
self.tp_rank,
)
self._previous_save_batch = None

def wait_for_save(self, connector_metadata: AscendConnectorMetadata) -> None:
current_event = None
assert self.kv_send_thread is not None
send_thread = self.kv_send_thread
requests: list[ReqMeta] = []

for request in connector_metadata.requests:
can_save = request.can_save
Expand All @@ -2414,11 +2438,16 @@ def wait_for_save(self, connector_metadata: AscendConnectorMetadata):
current_event.record()
request.skip_null_blocks_by_group = self.group_uses_align_state
request.current_event = current_event
send_thread.add_stored_request(request.req_id)
send_thread.add_request(request)
requests.append(request)

if current_event is not None:
send_thread.request_queue.join()
if not requests:
return

if not isinstance(send_thread, KVCacheStoreSendingThread):
raise TypeError(
f"Non-layerwise KV save requires KVCacheStoreSendingThread, but got {type(send_thread).__name__}"
)
self._previous_save_batch = send_thread.add_save_batch(requests)

def retrieve_layer(
self,
Expand Down
12 changes: 9 additions & 3 deletions vllm_ascend/worker/model_runner_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -2214,10 +2214,16 @@ def execute_model(
# scheduler_output in the worker, so no copy is needed.
pp_group = get_pp_group()
if pp_group.world_size > 1 and not pp_group.is_last_rank:
new_token_ids = scheduler_output.scheduled_cached_reqs.new_token_ids
cached_reqs = scheduler_output.scheduled_cached_reqs
new_token_ids = cached_reqs.new_token_ids
if new_token_ids and all(not token_ids for token_ids in new_token_ids):
scheduler_output = deepcopy(scheduler_output)
scheduler_output.scheduled_cached_reqs.new_token_ids = []
scheduler_output = replace(
scheduler_output,
scheduled_cached_reqs=replace(
cached_reqs,
new_token_ids=[],
),
)

if has_kv_transfer_group():
kv_connector_metadata = scheduler_output.kv_connector_metadata
Expand Down
Loading