Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
09cd90a196 | ||
|
|
b9685019bd | ||
|
|
e2ded9d884 | ||
|
|
d526f9d91f | ||
|
|
48a772a0d4 | ||
|
|
55f2c075cb | ||
|
|
97c97f6c1e | ||
|
|
1198bd0605 | ||
|
|
1f4aa13f6c | ||
|
|
f95ede55c5 |
@@ -40,6 +40,26 @@ from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv
|
||||
pytestmark = pytest.mark.cpu_test
|
||||
|
||||
|
||||
def _get_remote_waiting_queue(scheduler: Scheduler):
|
||||
return getattr(scheduler, "waiting_for_remote_kvs", None)
|
||||
|
||||
|
||||
def _num_waiting_requests(scheduler: Scheduler) -> int:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
|
||||
|
||||
|
||||
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
return list(remote_waiting)
|
||||
return [
|
||||
req
|
||||
for req in scheduler.waiting
|
||||
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
]
|
||||
|
||||
|
||||
def test_add_requests():
|
||||
scheduler = create_scheduler()
|
||||
requests = create_requests(num_requests=10)
|
||||
@@ -1120,7 +1140,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
|
||||
|
||||
# Requests should first transition to WAITING_FOR_REMOTE_KVS
|
||||
output = scheduler.schedule()
|
||||
assert len(scheduler.waiting) == len(req_ids)
|
||||
assert _num_waiting_requests(scheduler) == len(req_ids)
|
||||
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(output.scheduled_new_reqs) == 0
|
||||
for req in scheduler.requests.values():
|
||||
@@ -1139,7 +1160,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
|
||||
|
||||
# Simulate KV transfer completion using KVConnectorOutput.finished_recving
|
||||
output = scheduler.schedule()
|
||||
assert len(scheduler.waiting) == len(req_ids)
|
||||
assert _num_waiting_requests(scheduler) == len(req_ids)
|
||||
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
|
||||
assert len(scheduler.running) == 0
|
||||
|
||||
MODEL_RUNNER_OUTPUT = ModelRunnerOutput(
|
||||
@@ -1546,7 +1568,7 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
|
||||
# All can be scheduled - 1st token.
|
||||
output = scheduler.schedule()
|
||||
if is_async:
|
||||
assert len(scheduler.waiting) == 2
|
||||
assert _num_waiting_requests(scheduler) == 2
|
||||
assert scheduler.running == []
|
||||
_step_until_kv_transfer_finished(scheduler, req_ids)
|
||||
output = scheduler.schedule()
|
||||
@@ -1604,7 +1626,9 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
|
||||
# This will have a local and remote cache hit.
|
||||
output = scheduler.schedule()
|
||||
if is_async:
|
||||
waiting_req_ids = [req.request_id for req in scheduler.waiting]
|
||||
waiting_req_ids = [
|
||||
req.request_id for req in _get_remote_waiting_requests(scheduler)
|
||||
]
|
||||
assert len(waiting_req_ids) == 1
|
||||
_step_until_kv_transfer_finished(scheduler, waiting_req_ids)
|
||||
output = scheduler.schedule()
|
||||
@@ -3614,6 +3638,10 @@ def test_prepend_skipped_requests_order():
|
||||
# simulate first 2 waiting requests are waiting for remote KVs
|
||||
for req in expected_waiting_reqs[:2]:
|
||||
req.status = RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
scheduler.waiting.remove_request(req)
|
||||
remote_waiting.add_request(req)
|
||||
|
||||
# schedule step
|
||||
# expect the first 2 waiting to be skipped, the third running,
|
||||
@@ -3623,8 +3651,13 @@ def test_prepend_skipped_requests_order():
|
||||
# pop the third request which is expected to be running
|
||||
expected_waiting_reqs.pop(2)
|
||||
|
||||
# verify waiting order is preserved
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs
|
||||
# verify waiting order is preserved for schedulable requests.
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs[2:]
|
||||
assert list(remote_waiting) == expected_waiting_reqs[:2]
|
||||
else:
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs
|
||||
|
||||
|
||||
def test_abort_request_waiting_for_remote_kvs():
|
||||
|
||||
@@ -17,6 +17,26 @@ from .utils import (
|
||||
)
|
||||
|
||||
|
||||
def _get_remote_waiting_queue(scheduler: Scheduler):
|
||||
return getattr(scheduler, "waiting_for_remote_kvs", None)
|
||||
|
||||
|
||||
def _num_waiting_requests(scheduler: Scheduler) -> int:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
|
||||
|
||||
|
||||
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
return list(remote_waiting)
|
||||
return [
|
||||
req
|
||||
for req in scheduler.waiting
|
||||
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
]
|
||||
|
||||
|
||||
def _make_get_num_new_matched_tokens(
|
||||
req_num_new_matched_tokens: dict[str, int],
|
||||
async_load,
|
||||
@@ -76,8 +96,8 @@ def test_async_load_failure(
|
||||
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
assert len(scheduler.waiting) == 3
|
||||
for request in scheduler.waiting:
|
||||
assert _num_waiting_requests(scheduler) == 3
|
||||
for request in _get_remote_waiting_requests(scheduler):
|
||||
assert request.num_computed_tokens == 0
|
||||
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
assert scheduler.connector.get_num_new_matched_tokens.call_count == 3
|
||||
@@ -96,8 +116,8 @@ def test_async_load_failure(
|
||||
|
||||
min_invalid_block_idx = min(invalid_block_idxs)
|
||||
|
||||
assert len(scheduler.waiting) == 3
|
||||
for request in scheduler.waiting:
|
||||
assert _num_waiting_requests(scheduler) == 3
|
||||
for request in _get_remote_waiting_requests(scheduler):
|
||||
if request.request_id == request2.request_id:
|
||||
assert request.num_computed_tokens == (
|
||||
min_invalid_block_idx * scheduler.block_size
|
||||
@@ -303,8 +323,10 @@ def test_async_progressive_load_failure(
|
||||
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
assert len(scheduler.waiting) == 1
|
||||
assert scheduler.waiting.peek_request().request_id == request.request_id
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
|
||||
assert len(remote_waiting_reqs) == 1
|
||||
assert remote_waiting_reqs[0].request_id == request.request_id
|
||||
assert request.num_computed_tokens == 0
|
||||
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
assert scheduler.connector.get_num_new_matched_tokens.call_count == 1
|
||||
@@ -325,8 +347,10 @@ def test_async_progressive_load_failure(
|
||||
|
||||
min_invalid_block_idx = min(min_invalid_block_idx, invalid_block_idx)
|
||||
|
||||
assert len(scheduler.waiting) == 1
|
||||
assert scheduler.waiting.peek_request().request_id == request.request_id
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
|
||||
assert len(remote_waiting_reqs) == 1
|
||||
assert remote_waiting_reqs[0].request_id == request.request_id
|
||||
assert request.num_computed_tokens == (
|
||||
min_invalid_block_idx * scheduler.block_size
|
||||
)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from collections.abc import Iterable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
@@ -35,6 +36,61 @@ def async_copy_to_gpu(
|
||||
return out.copy_(tmp, non_blocking=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PrepareInputsRingSlot:
|
||||
idx_mapping_cpu: torch.Tensor
|
||||
query_start_loc_cpu: torch.Tensor
|
||||
cu_num_logits_cpu: torch.Tensor
|
||||
event: torch.cuda.Event
|
||||
in_use: bool = False
|
||||
|
||||
|
||||
class PrepareInputsBuffers:
|
||||
def __init__(self, ring_size: int, max_num_reqs: int, device: torch.device):
|
||||
self.device = device
|
||||
self.ring_slots = [
|
||||
PrepareInputsRingSlot(
|
||||
idx_mapping_cpu=torch.empty(
|
||||
max_num_reqs, dtype=torch.int32, pin_memory=True
|
||||
),
|
||||
query_start_loc_cpu=torch.empty(
|
||||
max_num_reqs + 1, dtype=torch.int32, pin_memory=True
|
||||
),
|
||||
cu_num_logits_cpu=torch.empty(
|
||||
max_num_reqs + 1, dtype=torch.int32, pin_memory=True
|
||||
),
|
||||
event=torch.cuda.Event(),
|
||||
)
|
||||
for _ in range(ring_size)
|
||||
]
|
||||
self.ring_slot_idx = -1
|
||||
self.idx_mapping_gpu = torch.empty(
|
||||
max_num_reqs, dtype=torch.int32, device=device
|
||||
)
|
||||
self.cu_num_logits_gpu = torch.empty(
|
||||
max_num_reqs + 1, dtype=torch.int32, device=device
|
||||
)
|
||||
self.arange_reqs_np = np.arange(max_num_reqs + 1, dtype=np.int32)
|
||||
self.arange_reqs_gpu = torch.arange(
|
||||
max_num_reqs + 1, dtype=torch.int32, device=device
|
||||
)
|
||||
self.zero_local_pos_gpu = torch.zeros(
|
||||
max_num_reqs, dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
def acquire_ring_slot(self) -> PrepareInputsRingSlot:
|
||||
slot_idx = (self.ring_slot_idx + 1) % len(self.ring_slots)
|
||||
self.ring_slot_idx = slot_idx
|
||||
slot = self.ring_slots[slot_idx]
|
||||
if slot.in_use and not slot.event.query():
|
||||
slot.event.synchronize()
|
||||
return slot
|
||||
|
||||
def mark_ring_slot_inflight(self, slot: PrepareInputsRingSlot) -> None:
|
||||
slot.event.record(torch.cuda.current_stream(self.device))
|
||||
slot.in_use = True
|
||||
|
||||
|
||||
class UvaBuffer:
|
||||
def __init__(self, size: int | Sequence[int], dtype: torch.dtype):
|
||||
if not is_uva_available():
|
||||
|
||||
@@ -53,7 +53,7 @@ from vllm.v1.worker.gpu.attn_utils import (
|
||||
init_kv_cache,
|
||||
)
|
||||
from vllm.v1.worker.gpu.block_table import BlockTables
|
||||
from vllm.v1.worker.gpu.buffer_utils import async_copy_to_gpu
|
||||
from vllm.v1.worker.gpu.buffer_utils import PrepareInputsBuffers
|
||||
from vllm.v1.worker.gpu.cp_utils import prepare_dcp_local_seq_lens
|
||||
from vllm.v1.worker.gpu.cudagraph_utils import CudaGraphManager
|
||||
from vllm.v1.worker.gpu.dp_utils import (
|
||||
@@ -187,6 +187,12 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
max_num_tokens=self.max_num_tokens,
|
||||
device=self.device,
|
||||
)
|
||||
# ring-buffered staging for prepare_inputs
|
||||
self._prepare_inputs_buffers = PrepareInputsBuffers(
|
||||
ring_size=max(self.pp_size, 2),
|
||||
max_num_reqs=self.max_num_reqs,
|
||||
device=self.device,
|
||||
)
|
||||
self.sampler = Sampler(
|
||||
max_num_reqs=self.max_num_reqs,
|
||||
vocab_size=self.vocab_size,
|
||||
@@ -567,6 +573,8 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
assert num_tokens > 0
|
||||
num_tokens_per_req = scheduler_output.num_scheduled_tokens
|
||||
num_reqs = len(num_tokens_per_req)
|
||||
prep_buffers = self._prepare_inputs_buffers
|
||||
slot = prep_buffers.acquire_ring_slot()
|
||||
|
||||
# Decode first, then prefill.
|
||||
# batch_idx -> req_id
|
||||
@@ -576,7 +584,10 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
|
||||
idx_mapping_iter = map(self.req_states.req_id_to_index.get, req_ids)
|
||||
idx_mapping_np = np.fromiter(idx_mapping_iter, dtype=np.int32, count=num_reqs)
|
||||
idx_mapping = async_copy_to_gpu(idx_mapping_np, device=self.device)
|
||||
idx_mapping_cpu = slot.idx_mapping_cpu[:num_reqs]
|
||||
np.copyto(idx_mapping_cpu.numpy(), idx_mapping_np, casting="no")
|
||||
idx_mapping = prep_buffers.idx_mapping_gpu[:num_reqs]
|
||||
idx_mapping.copy_(idx_mapping_cpu, non_blocking=True)
|
||||
|
||||
# Get the number of draft tokens for each request.
|
||||
draft_tokens = scheduler_output.scheduled_spec_decode_tokens
|
||||
@@ -584,14 +595,10 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
# No draft token scheduled (common case).
|
||||
total_num_draft_tokens = 0
|
||||
total_num_logits = num_reqs
|
||||
cu_num_logits_np = np.arange(num_reqs + 1, dtype=np.int32)
|
||||
cu_num_logits = torch.arange(
|
||||
num_reqs + 1, device=self.device, dtype=torch.int32
|
||||
)
|
||||
cu_num_logits_np = prep_buffers.arange_reqs_np[: num_reqs + 1]
|
||||
cu_num_logits = prep_buffers.arange_reqs_gpu[: num_reqs + 1]
|
||||
expanded_idx_mapping = idx_mapping
|
||||
expanded_local_pos = torch.zeros(
|
||||
num_reqs, dtype=torch.int32, device=self.device
|
||||
)
|
||||
expanded_local_pos = prep_buffers.zero_local_pos_gpu[:num_reqs]
|
||||
else:
|
||||
num_draft_tokens = np.array(
|
||||
[len(draft_tokens.get(req_id, ())) for req_id in req_ids],
|
||||
@@ -601,10 +608,14 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
total_num_logits = num_reqs + total_num_draft_tokens
|
||||
|
||||
num_logits = num_draft_tokens + 1
|
||||
cu_num_logits_np = np.empty(num_reqs + 1, dtype=np.int32)
|
||||
cu_num_logits_np[0] = 0
|
||||
np.cumsum(num_logits, out=cu_num_logits_np[1:])
|
||||
cu_num_logits = async_copy_to_gpu(cu_num_logits_np, device=self.device)
|
||||
cu_num_logits_cpu = slot.cu_num_logits_cpu[: num_reqs + 1]
|
||||
cu_num_logits_cpu_np = cu_num_logits_cpu.numpy()
|
||||
cu_num_logits_cpu_np[0] = 0
|
||||
np.cumsum(num_logits, out=cu_num_logits_cpu_np[1:])
|
||||
cu_num_logits = prep_buffers.cu_num_logits_gpu[: num_reqs + 1]
|
||||
cu_num_logits.copy_(cu_num_logits_cpu, non_blocking=True)
|
||||
# keep an independent CPU snapshot because ring slots are reused.
|
||||
cu_num_logits_np = cu_num_logits_cpu_np.copy()
|
||||
|
||||
max_expand_len = self.num_speculative_steps + 1
|
||||
expanded_idx_mapping, expanded_local_pos = expand_idx_mapping(
|
||||
@@ -612,14 +623,20 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
)
|
||||
|
||||
# Get query_start_loc.
|
||||
query_start_loc_np = np.empty(self.max_num_reqs + 1, dtype=np.int32)
|
||||
query_start_loc_np[0] = 0
|
||||
np.cumsum(num_scheduled_tokens, out=query_start_loc_np[1 : num_reqs + 1])
|
||||
query_start_loc_cpu_full = slot.query_start_loc_cpu
|
||||
query_start_loc_np_full = query_start_loc_cpu_full.numpy()
|
||||
query_start_loc_np_full[0] = 0
|
||||
np.cumsum(num_scheduled_tokens, out=query_start_loc_np_full[1 : num_reqs + 1])
|
||||
# Pad for full CUDA graph mode.
|
||||
# Some attention backends like FA3 require query_start_loc to be non-decreasing.
|
||||
query_start_loc_np[num_reqs + 1 :] = num_tokens
|
||||
async_copy_to_gpu(query_start_loc_np, out=self.input_buffers.query_start_loc)
|
||||
query_start_loc_np = query_start_loc_np[: num_reqs + 1]
|
||||
query_start_loc_np_full[num_reqs + 1 :] = num_tokens
|
||||
self.input_buffers.query_start_loc.copy_(
|
||||
query_start_loc_cpu_full, non_blocking=True
|
||||
)
|
||||
prep_buffers.mark_ring_slot_inflight(slot)
|
||||
|
||||
# keep an independent CPU snapshot because ring slots are reused.
|
||||
query_start_loc_np = query_start_loc_np_full[: num_reqs + 1].copy()
|
||||
query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
|
||||
|
||||
# Get prefill tokens if any.
|
||||
|
||||
Reference in New Issue
Block a user