Compare commits

...
Author SHA1 Message Date
yewentao256 09cd90a196 fix unit test
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-03-02 22:11:52 +00:00
yewentao256 b9685019bd Merge branch 'main' into wentao-optimize-model-runner-v2-prepare_inputs
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-03-02 21:00:46 +00:00
yewentao256 e2ded9d884 Revert "remove ring buffer"
This reverts commit 48a772a0d4.
2026-03-02 20:59:14 +00:00
Wentao YeandGitHub d526f9d91f Merge branch 'main' into wentao-optimize-model-runner-v2-prepare_inputs 2026-02-27 13:28:10 -05:00
yewentao256 48a772a0d4 remove ring buffer
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-02-27 18:26:11 +00:00
yewentao256 55f2c075cb merge main and refactor
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-02-26 20:32:06 +00:00
yewentao256 97c97f6c1e Merge branch 'main' into wentao-optimize-model-runner-v2-prepare_inputs
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-02-26 19:16:34 +00:00
yewentao256 1198bd0605 add a ring slot to make sure no race
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-02-26 19:11:20 +00:00
yewentao256 1f4aa13f6c Merge branch 'main' into wentao-optimize-model-runner-v2-prepare_inputs 2026-02-26 15:50:28 +00:00
yewentao256 f95ede55c5 optimize model runner v2 prepare_inputs by reducing memory alloc
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-02-25 21:54:22 +00:00
4 changed files with 163 additions and 33 deletions
+39 -6
View File
@@ -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
)
+56
View File
@@ -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():
+36 -19
View File
@@ -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.