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 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(): def test_add_requests():
scheduler = create_scheduler() scheduler = create_scheduler()
requests = create_requests(num_requests=10) 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 # Requests should first transition to WAITING_FOR_REMOTE_KVS
output = scheduler.schedule() 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(scheduler.running) == 0
assert len(output.scheduled_new_reqs) == 0 assert len(output.scheduled_new_reqs) == 0
for req in scheduler.requests.values(): 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 # Simulate KV transfer completion using KVConnectorOutput.finished_recving
output = scheduler.schedule() 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(scheduler.running) == 0
MODEL_RUNNER_OUTPUT = ModelRunnerOutput( 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. # All can be scheduled - 1st token.
output = scheduler.schedule() output = scheduler.schedule()
if is_async: if is_async:
assert len(scheduler.waiting) == 2 assert _num_waiting_requests(scheduler) == 2
assert scheduler.running == [] assert scheduler.running == []
_step_until_kv_transfer_finished(scheduler, req_ids) _step_until_kv_transfer_finished(scheduler, req_ids)
output = scheduler.schedule() 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. # This will have a local and remote cache hit.
output = scheduler.schedule() output = scheduler.schedule()
if is_async: 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 assert len(waiting_req_ids) == 1
_step_until_kv_transfer_finished(scheduler, waiting_req_ids) _step_until_kv_transfer_finished(scheduler, waiting_req_ids)
output = scheduler.schedule() output = scheduler.schedule()
@@ -3614,6 +3638,10 @@ def test_prepend_skipped_requests_order():
# simulate first 2 waiting requests are waiting for remote KVs # simulate first 2 waiting requests are waiting for remote KVs
for req in expected_waiting_reqs[:2]: for req in expected_waiting_reqs[:2]:
req.status = RequestStatus.WAITING_FOR_REMOTE_KVS 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 # schedule step
# expect the first 2 waiting to be skipped, the third running, # 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 # pop the third request which is expected to be running
expected_waiting_reqs.pop(2) expected_waiting_reqs.pop(2)
# verify waiting order is preserved # verify waiting order is preserved for schedulable requests.
assert list(scheduler.waiting) == expected_waiting_reqs 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(): 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( def _make_get_num_new_matched_tokens(
req_num_new_matched_tokens: dict[str, int], req_num_new_matched_tokens: dict[str, int],
async_load, async_load,
@@ -76,8 +96,8 @@ def test_async_load_failure(
scheduler_output = scheduler.schedule() scheduler_output = scheduler.schedule()
assert len(scheduler.waiting) == 3 assert _num_waiting_requests(scheduler) == 3
for request in scheduler.waiting: for request in _get_remote_waiting_requests(scheduler):
assert request.num_computed_tokens == 0 assert request.num_computed_tokens == 0
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert scheduler.connector.get_num_new_matched_tokens.call_count == 3 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) min_invalid_block_idx = min(invalid_block_idxs)
assert len(scheduler.waiting) == 3 assert _num_waiting_requests(scheduler) == 3
for request in scheduler.waiting: for request in _get_remote_waiting_requests(scheduler):
if request.request_id == request2.request_id: if request.request_id == request2.request_id:
assert request.num_computed_tokens == ( assert request.num_computed_tokens == (
min_invalid_block_idx * scheduler.block_size min_invalid_block_idx * scheduler.block_size
@@ -303,8 +323,10 @@ def test_async_progressive_load_failure(
scheduler_output = scheduler.schedule() scheduler_output = scheduler.schedule()
assert len(scheduler.waiting) == 1 assert _num_waiting_requests(scheduler) == 1
assert scheduler.waiting.peek_request().request_id == request.request_id 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.num_computed_tokens == 0
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
assert scheduler.connector.get_num_new_matched_tokens.call_count == 1 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) min_invalid_block_idx = min(min_invalid_block_idx, invalid_block_idx)
assert len(scheduler.waiting) == 1 assert _num_waiting_requests(scheduler) == 1
assert scheduler.waiting.peek_request().request_id == request.request_id 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 == ( assert request.num_computed_tokens == (
min_invalid_block_idx * scheduler.block_size min_invalid_block_idx * scheduler.block_size
) )
+56
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project # SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Iterable, Sequence from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from functools import partial from functools import partial
import numpy as np import numpy as np
@@ -35,6 +36,61 @@ def async_copy_to_gpu(
return out.copy_(tmp, non_blocking=True) 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: class UvaBuffer:
def __init__(self, size: int | Sequence[int], dtype: torch.dtype): def __init__(self, size: int | Sequence[int], dtype: torch.dtype):
if not is_uva_available(): if not is_uva_available():
+36 -19
View File
@@ -53,7 +53,7 @@ from vllm.v1.worker.gpu.attn_utils import (
init_kv_cache, init_kv_cache,
) )
from vllm.v1.worker.gpu.block_table import BlockTables 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.cp_utils import prepare_dcp_local_seq_lens
from vllm.v1.worker.gpu.cudagraph_utils import CudaGraphManager from vllm.v1.worker.gpu.cudagraph_utils import CudaGraphManager
from vllm.v1.worker.gpu.dp_utils import ( from vllm.v1.worker.gpu.dp_utils import (
@@ -187,6 +187,12 @@ class GPUModelRunner(LoRAModelRunnerMixin):
max_num_tokens=self.max_num_tokens, max_num_tokens=self.max_num_tokens,
device=self.device, 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( self.sampler = Sampler(
max_num_reqs=self.max_num_reqs, max_num_reqs=self.max_num_reqs,
vocab_size=self.vocab_size, vocab_size=self.vocab_size,
@@ -567,6 +573,8 @@ class GPUModelRunner(LoRAModelRunnerMixin):
assert num_tokens > 0 assert num_tokens > 0
num_tokens_per_req = scheduler_output.num_scheduled_tokens num_tokens_per_req = scheduler_output.num_scheduled_tokens
num_reqs = len(num_tokens_per_req) num_reqs = len(num_tokens_per_req)
prep_buffers = self._prepare_inputs_buffers
slot = prep_buffers.acquire_ring_slot()
# Decode first, then prefill. # Decode first, then prefill.
# batch_idx -> req_id # 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_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_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. # Get the number of draft tokens for each request.
draft_tokens = scheduler_output.scheduled_spec_decode_tokens draft_tokens = scheduler_output.scheduled_spec_decode_tokens
@@ -584,14 +595,10 @@ class GPUModelRunner(LoRAModelRunnerMixin):
# No draft token scheduled (common case). # No draft token scheduled (common case).
total_num_draft_tokens = 0 total_num_draft_tokens = 0
total_num_logits = num_reqs total_num_logits = num_reqs
cu_num_logits_np = np.arange(num_reqs + 1, dtype=np.int32) cu_num_logits_np = prep_buffers.arange_reqs_np[: num_reqs + 1]
cu_num_logits = torch.arange( cu_num_logits = prep_buffers.arange_reqs_gpu[: num_reqs + 1]
num_reqs + 1, device=self.device, dtype=torch.int32
)
expanded_idx_mapping = idx_mapping expanded_idx_mapping = idx_mapping
expanded_local_pos = torch.zeros( expanded_local_pos = prep_buffers.zero_local_pos_gpu[:num_reqs]
num_reqs, dtype=torch.int32, device=self.device
)
else: else:
num_draft_tokens = np.array( num_draft_tokens = np.array(
[len(draft_tokens.get(req_id, ())) for req_id in req_ids], [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 total_num_logits = num_reqs + total_num_draft_tokens
num_logits = num_draft_tokens + 1 num_logits = num_draft_tokens + 1
cu_num_logits_np = np.empty(num_reqs + 1, dtype=np.int32) cu_num_logits_cpu = slot.cu_num_logits_cpu[: num_reqs + 1]
cu_num_logits_np[0] = 0 cu_num_logits_cpu_np = cu_num_logits_cpu.numpy()
np.cumsum(num_logits, out=cu_num_logits_np[1:]) cu_num_logits_cpu_np[0] = 0
cu_num_logits = async_copy_to_gpu(cu_num_logits_np, device=self.device) 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 max_expand_len = self.num_speculative_steps + 1
expanded_idx_mapping, expanded_local_pos = expand_idx_mapping( expanded_idx_mapping, expanded_local_pos = expand_idx_mapping(
@@ -612,14 +623,20 @@ class GPUModelRunner(LoRAModelRunnerMixin):
) )
# Get query_start_loc. # Get query_start_loc.
query_start_loc_np = np.empty(self.max_num_reqs + 1, dtype=np.int32) query_start_loc_cpu_full = slot.query_start_loc_cpu
query_start_loc_np[0] = 0 query_start_loc_np_full = query_start_loc_cpu_full.numpy()
np.cumsum(num_scheduled_tokens, out=query_start_loc_np[1 : num_reqs + 1]) 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. # Pad for full CUDA graph mode.
# Some attention backends like FA3 require query_start_loc to be non-decreasing. # Some attention backends like FA3 require query_start_loc to be non-decreasing.
query_start_loc_np[num_reqs + 1 :] = num_tokens query_start_loc_np_full[num_reqs + 1 :] = num_tokens
async_copy_to_gpu(query_start_loc_np, out=self.input_buffers.query_start_loc) self.input_buffers.query_start_loc.copy_(
query_start_loc_np = query_start_loc_np[: num_reqs + 1] 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] query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
# Get prefill tokens if any. # Get prefill tokens if any.