forked from Karylab-cklius/vllm
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
|
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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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():
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user