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 @@ -2846,6 +2846,46 @@ def test_pending_transfer_defers_prefix_lookup():
scheduler.manager.lookup.assert_not_called()


@pytest.mark.parametrize("async_scheduling", [False, True])
def test_prefix_cache_reset_preserves_same_step_offload_reload(
request_runner, async_scheduling: bool
):
"""A new load after reset must survive the same batch's preemption notice."""
runner = request_runner(
block_size=4, num_gpu_blocks=100, async_scheduling=async_scheduling
)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.new_request(token_ids=[0] * 8)
runner.run(decoded_tokens=[0, 0], expected_stored=(0, 1))

request = runner.scheduler.requests["0"]
req_status = runner.connector_scheduler._req_status["0"]
assert request.status == RequestStatus.RUNNING
assert not req_status.transfer_jobs
assert runner.scheduler.reset_prefix_cache(reset_running_requests=True)
assert request.status == RequestStatus.PREEMPTED

runner.connector_scheduler._maximal_prefix_lookup = lambda keys, ctx, *_: 2
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output([])
)
runner.run(decoded_tokens=[], complete_transfers=False)

assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert len(req_status.transfer_jobs) == 1
(job_id,) = req_status.transfer_jobs
assert not runner.connector_scheduler._jobs[job_id].is_store
assert job_id in runner.offloading_spec.handler.waiting_jobs
runner.manager.complete_load.assert_not_called()

runner.run(decoded_tokens=[EOS_TOKEN_ID], expected_loaded=(0, 1))
runner.manager.complete_load.assert_called_once()
assert not runner.connector_scheduler._jobs
assert not runner.scheduler.requests


def test_async_preempt_readmit_before_transfer_output_is_deferred(request_runner):
"""A preempted request can be scheduled again before flush output is read.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1856,11 +1856,18 @@ def build_connector_meta(
# Flush jobs for preempted requests.
for req_id in scheduler_output.preempted_req_ids or ():
req_status = self._req_status.get(req_id)
if req_status is None or not req_status.transfer_jobs:
if req_status is None:
continue
# A cache reset can preempt and reload a request in the same step.
# These new loads have not been submitted to the worker yet.
jobs_to_flush = (
req_status.transfer_jobs - self._current_batch_load_jobs.keys()
)
if not jobs_to_flush:
continue
any_jid = next(iter(req_status.transfer_jobs))
any_jid = next(iter(jobs_to_flush))
assert self._jobs[any_jid].is_store
self._current_batch_jobs_to_flush.update(req_status.transfer_jobs)
self._current_batch_jobs_to_flush.update(jobs_to_flush)

# Flush jobs that contain re-allocated blocks.
if (
Expand Down
Loading