Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -2478,6 +2478,42 @@ def test_reset_cache(request_runner, async_scheduling: bool):
assert group_state.next_stored_chunk_idx == 0


def test_reset_cache_flush_is_delivered_when_idle_and_reset_is_reentrant(
request_runner,
):
"""A reset that discards an in-flight store after the last request
finished must keep the engine stepping until its flush set reaches the
workers, and a second reset before that step must not assert (RL loops
call reset_prefix_cache(reset_connector=True) every iteration, and
EngineCore drains all queued utility calls before it steps again)."""
runner = request_runner(
block_size=4, num_gpu_blocks=100, async_scheduling=False, blocks_per_chunk=1
)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.manager.has_pending_work.return_value = False
runner.new_request(token_ids=[0] * 8)
runner.run(decoded_tokens=[EOS_TOKEN_ID], complete_transfers=False)
cs = runner.connector_scheduler
store_job_ids = set(cs._jobs)
assert store_job_ids
assert not runner.scheduler.has_unfinished_requests()

cs.reset_cache()
cs.reset_cache()

assert cs._current_batch_jobs_to_flush == store_job_ids
assert cs.has_pending_push_work()
assert runner.scheduler.has_requests()

scheduler_output = runner.scheduler.schedule()
meta = scheduler_output.kv_connector_metadata
assert isinstance(meta, OffloadingConnectorMetadata)
assert meta.jobs_to_flush == store_job_ids
assert not cs.has_pending_push_work()


@pytest.mark.parametrize("async_scheduling", [True, False])
def test_reset_cache_finalizes_finished_request_with_pending_store(
request_runner, async_scheduling: bool
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1835,8 +1835,14 @@ def has_pending_push_work(self) -> bool:

While True, build_connector_meta() and update_connector_output()
continue to be called even when no requests are scheduled.
A flush set left by reset_cache() counts: it reaches the workers only
through build_connector_meta().
"""
return bool(self._jobs) or self.manager.has_pending_work()
return (
bool(self._jobs)
or bool(self._current_batch_jobs_to_flush)
or self.manager.has_pending_work()
)

def update_connector_output(self, connector_output: KVConnectorOutput):
"""Update KVConnector state from worker-side connectors output.
Expand Down Expand Up @@ -1990,9 +1996,10 @@ def take_events(self) -> Iterable[KVCacheEvent]:

def reset_cache(self) -> None:
"""Reset the offloading manager cache, evicting all stored chunks."""
# reset_cache cannot be called in the middle of a schedule step
# reset_cache cannot be called in the middle of a schedule step.
# _current_batch_jobs_to_flush may still hold a previous reset's flush
# set if no step ran since; the new ids are merged into it.
assert not self._current_batch_load_jobs
assert not self._current_batch_jobs_to_flush
assert not self._current_batch_allocated_block_ids

# Flush all in-flight jobs
Expand Down
Loading