Skip to content

Commit d0ce57a

Browse files
Stage mixed HiSparse host imports through GPU
Route same-node host-bound HiSparse regions through the generic NIXL host stager while preserving direct GPU transfers for device-resident regions. Complete mixed receives only after both transfer paths finish and preserve failure reporting across sibling cancellation. Co-authored-by: OpenAI Codex <codex@openai.com> Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
1 parent ded04a1 commit d0ce57a

4 files changed

Lines changed: 233 additions & 28 deletions

File tree

tests/v1/kv_connector/unit/test_nixl_connector.py

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
import os
77
import tempfile
88
import textwrap
9+
import threading
910
import time
1011
import uuid
1112
from collections import defaultdict
@@ -1906,6 +1907,67 @@ def test_host_stager_is_only_initialized_for_same_host_reads(monkeypatch):
19061907
stager_factory.assert_called_once()
19071908

19081909

1910+
@pytest.mark.cpu_test
1911+
@pytest.mark.parametrize("staging_finishes_first", [False, True])
1912+
def test_mixed_host_staging_waits_for_both_receive_parts(staging_finishes_first):
1913+
"""A mixed host/device receive completes only after both paths finish."""
1914+
worker = object.__new__(NixlConnectorWorker)
1915+
worker._recving_metadata = {"request": MagicMock()}
1916+
worker._recving_transfers = {"request": [7]}
1917+
worker._failed_recv_lock = threading.Lock()
1918+
worker._failed_recv_pending = set()
1919+
worker._failed_recv_reported = set()
1920+
worker._pending_recv_notifs = {"request": [("agent", b"notification")]}
1921+
worker._send_pending_recv_notifs = MagicMock()
1922+
worker._report_failed_recv = MagicMock()
1923+
worker.xfer_stats = MagicMock()
1924+
worker.nixl_wrapper = MagicMock()
1925+
worker.nixl_wrapper.check_xfer_state.return_value = "DONE"
1926+
1927+
stager = MagicMock()
1928+
stager.active_req_ids = {"request"}
1929+
1930+
def finish_staging():
1931+
stager.active_req_ids = set()
1932+
return {"request"}, set()
1933+
1934+
stager.get_finished.side_effect = finish_staging
1935+
worker._host_stager = stager
1936+
1937+
if staging_finishes_first:
1938+
assert worker._get_finished_host_staging() == set()
1939+
assert worker._pop_done_transfers(worker._recving_transfers, is_recv=True) == {
1940+
"request"
1941+
}
1942+
else:
1943+
assert (
1944+
worker._pop_done_transfers(worker._recving_transfers, is_recv=True) == set()
1945+
)
1946+
assert worker._get_finished_host_staging() == {"request"}
1947+
1948+
worker._send_pending_recv_notifs.assert_called_once_with("request")
1949+
1950+
1951+
@pytest.mark.cpu_test
1952+
def test_mixed_host_staging_reports_failure_after_aborted_sibling_drains():
1953+
worker = object.__new__(NixlConnectorWorker)
1954+
worker._recving_metadata = {"request": MagicMock()}
1955+
worker._recving_transfers = {}
1956+
worker._failed_recv_lock = threading.Lock()
1957+
worker._failed_recv_pending = {"request"}
1958+
worker._pending_recv_notifs = {}
1959+
worker._report_failed_recv = MagicMock()
1960+
1961+
stager = MagicMock()
1962+
stager.active_req_ids = set()
1963+
stager.get_finished.return_value = set(), set()
1964+
worker._host_stager = stager
1965+
1966+
assert worker._get_finished_host_staging() == set()
1967+
assert worker._failed_recv_pending == set()
1968+
worker._report_failed_recv.assert_called_once_with("request")
1969+
1970+
19091971
def _run_abort_timeout_test(llm: LLM, timeout: int):
19101972
"""Helper function to run the abort timeout test logic."""
19111973
remote_prefill_opts = {

tests/v1/kv_connector/unit/test_nixl_connector_hma.py

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,18 +3,25 @@
33
"""Unit tests for NixlConnectorScheduler with HMA and Mamba N-1 prefill."""
44

55
import gc
6+
from types import SimpleNamespace
67
from unittest.mock import MagicMock, patch
78

89
import msgspec
10+
import numpy as np
911
import pytest
1012
import torch
1113

1214
from tests.v1.attention.utils import MockMambaBuilder
1315
from vllm import LLM, SamplingParams
1416
from vllm.config import KVTransferConfig, set_current_vllm_config
17+
from vllm.distributed.kv_transfer.kv_connector.v1.hisparse.nixl import (
18+
HiSparseNixlDestination,
19+
)
1520
from vllm.distributed.kv_transfer.kv_connector.v1.nixl import base_worker as bw
1621
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
1722
NixlAgentMetadata,
23+
RemoteMeta,
24+
ReqMeta,
1825
)
1926
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
2027
NixlConnectorWorker,
@@ -1971,6 +1978,83 @@ def test_hisparse_host_import_keeps_host_blocks_out_of_gpu_regions():
19711978
assert request_metadata.local_block_ids == ([10, 11], [20])
19721979

19731980

1981+
@pytest.mark.cpu_test
1982+
def test_hisparse_same_host_import_stages_host_regions():
1983+
"""Host regions use GPU staging while device-only regions transfer directly."""
1984+
destination = object.__new__(HiSparseNixlDestination)
1985+
destination.host_regions = [(1_000, 16), None]
1986+
destination._descriptor_offsets = [0, None]
1987+
destination._xfer_handle = 100
1988+
destination._host_desc_lens = np.full(4, 16, dtype=np.int64)
1989+
destination._host_addrs = np.arange(1_000, 1_064, 16, dtype=np.uint64)
1990+
host_buffer = torch.empty(64, dtype=torch.uint8)
1991+
destination._host_pools = {host_buffer.data_ptr(): (host_buffer, [])}
1992+
1993+
worker = MagicMock()
1994+
worker._has_mamba = False
1995+
worker.use_mla = True
1996+
worker.transfer_topo = MagicMock()
1997+
worker.transfer_topo.get_engine_info.return_value = SimpleNamespace(
1998+
remote_physical_blocks_per_logical=1,
1999+
remote_block_size=64,
2000+
)
2001+
worker.transfer_topo.tp_ratio.return_value = 1
2002+
worker._block_ids_by_region.side_effect = [
2003+
[[1, 2], [1, 2]],
2004+
[[3, 4], [5, 6]],
2005+
]
2006+
worker._apply_prefix_caching_by_region.side_effect = lambda local, remote: (
2007+
local,
2008+
remote,
2009+
)
2010+
worker._compute_desc_ids.side_effect = [
2011+
np.array([10, 11, 20, 21]),
2012+
np.array([30, 31]),
2013+
]
2014+
worker.dst_num_blocks = {"remote": 8, "local": 8}
2015+
worker.dst_region_num_blocks = {"remote": [8, 8]}
2016+
worker.region_num_blocks = [8, 8]
2017+
worker.region_group_ids = [0, 1]
2018+
worker.num_regions = 2
2019+
worker.engine_id = "local"
2020+
worker._physical_blocks_per_logical_kv_block = 1
2021+
worker.src_xfer_handles_by_block_size = {64: 200}
2022+
worker.dst_xfer_side_handles = {"remote": {0: 300}}
2023+
worker._remote_agents = {"remote": {(0, 0): "agent"}}
2024+
worker._pending_recv_notifs = {}
2025+
worker._recving_transfers = {}
2026+
worker.nixl_wrapper.make_prepped_xfer.return_value = 400
2027+
stager = MagicMock()
2028+
worker._maybe_init_host_stager_for_buffers.return_value = stager
2029+
2030+
meta = ReqMeta(
2031+
local_block_ids=([3, 4], [5, 6]),
2032+
local_physical_block_ids=([3, 4], [5, 6]),
2033+
tp_size=1,
2034+
remote=RemoteMeta(
2035+
block_ids=([1, 2], [1, 2]),
2036+
host="local-host",
2037+
port=1234,
2038+
engine_id="remote",
2039+
request_id="remote-request",
2040+
),
2041+
hisparse_host_block_ids=[0, 1],
2042+
)
2043+
plan = MagicMock(all_source_ranks=(0,), local_consumers=1)
2044+
2045+
destination.read_host_blocks(worker, "request", meta, plan, [0, 1])
2046+
2047+
stager.submit.assert_called_once()
2048+
submit_args = stager.submit.call_args.args
2049+
assert submit_args[0] == "request"
2050+
np.testing.assert_array_equal(submit_args[1], [10, 11])
2051+
np.testing.assert_array_equal(submit_args[2], [0, 1])
2052+
assert submit_args[3] == 300
2053+
worker.nixl_wrapper.make_prepped_xfer.assert_called_once()
2054+
worker.nixl_wrapper.transfer.assert_called_once_with(400)
2055+
assert worker._recving_transfers == {"request": [400]}
2056+
2057+
19742058
@pytest.mark.cpu_test
19752059
def test_register_kv_caches_hybrid_mla_dual_purpose_regions():
19762060
"""Hybrid MLA+KDA registration: HMA tensors shared by both layer types

vllm/distributed/kv_transfer/kv_connector/v1/hisparse/nixl.py

Lines changed: 28 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -33,12 +33,16 @@ def __init__(self, kv_cache_config: KVCacheConfig, vllm_config: VllmConfig) -> N
3333
self._forward_context = vllm_config.compilation_config.static_forward_context
3434
self.host_regions: list[tuple[int, int] | None] = []
3535
self._host_pools: dict[int, tuple[torch.Tensor, list[HiSparseRuntime]]] = {}
36+
self._host_desc_lens = np.empty(0, dtype=np.int64)
37+
self._host_addrs = np.empty(0, dtype=np.uint64)
3638
self._descriptor_offsets: list[int | None] = []
3739
self._xfer_handle: int | None = None
3840

3941
def reset_regions(self) -> None:
4042
self.host_regions.clear()
4143
self._host_pools.clear()
44+
self._host_desc_lens = np.empty(0, dtype=np.int64)
45+
self._host_addrs = np.empty(0, dtype=np.uint64)
4246
self._descriptor_offsets.clear()
4347
self._xfer_handle = None
4448

@@ -140,6 +144,10 @@ def prepare_host_descriptors(
140144
for block_id in range(self.host_num_blocks)
141145
)
142146
descs = worker.nixl_wrapper.get_xfer_descs(blocks, "DRAM")
147+
self._host_addrs = np.asarray([block[0] for block in blocks], dtype=np.uint64)
148+
self._host_desc_lens = np.asarray(
149+
[block[1] for block in blocks], dtype=np.int64
150+
)
143151
self._xfer_handle = worker.nixl_wrapper.prep_xfer_dlist(
144152
"NIXL_INIT_AGENT", descs, backends=worker.nixl_backends
145153
)
@@ -242,18 +250,25 @@ def read_host_blocks(
242250
local_device_handle = worker.src_xfer_handles_by_block_size[
243251
remote_info.remote_block_size
244252
]
253+
host_local_ids_array = np.asarray(host_local_ids, dtype=np.int64)
254+
host_remote_ids_array = np.asarray(host_remote_ids, dtype=np.int64)
255+
host_stager = worker._maybe_init_host_stager_for_buffers(
256+
meta.remote.host,
257+
self._host_desc_lens,
258+
self._host_addrs,
259+
[pool for pool, _ in self._host_pools.values()],
260+
)
245261
transfer_specs = [
246-
(
247-
self._xfer_handle,
248-
np.asarray(host_local_ids, dtype=np.int64),
249-
np.asarray(host_remote_ids, dtype=np.int64),
250-
),
251262
(
252263
local_device_handle,
253264
device_local_ids,
254265
np.asarray(device_remote_ids, dtype=np.int64),
255-
),
266+
)
256267
]
268+
if host_stager is None:
269+
transfer_specs.insert(
270+
0, (self._xfer_handle, host_local_ids_array, host_remote_ids_array)
271+
)
257272

258273
notif_id = f"{meta.remote.request_id}:{plan.local_consumers}".encode()
259274
agents = worker._remote_agents[engine_id]
@@ -278,6 +293,13 @@ def read_host_blocks(
278293
selected_remote_ids,
279294
)
280295
)
296+
if host_stager is not None and len(host_local_ids_array):
297+
host_stager.submit(
298+
request_id,
299+
host_remote_ids_array,
300+
host_local_ids_array,
301+
remote_handle,
302+
)
281303
worker._recving_transfers.setdefault(request_id, [])
282304
for handle in handles:
283305
worker.nixl_wrapper.transfer(handle)

vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py

Lines changed: 59 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -2814,25 +2814,42 @@ def _handle_heartbeat(self, payload: str) -> None:
28142814

28152815
def _maybe_init_host_stager(self, remote_host: str) -> HostWriteStager | None:
28162816
"""Enable device staging for same-host host-buffer reads."""
2817-
if remote_host != envs.VLLM_NIXL_SIDE_CHANNEL_HOST:
2818-
return None
2819-
if self._host_stager is not None or self._host_stager_init_attempted:
2820-
return self._host_stager
2821-
stage_bytes = envs.VLLM_NIXL_HOST_STAGE_BYTES
2822-
if stage_bytes <= 0 or not self.use_host_buffer:
2817+
if not self.use_host_buffer:
28232818
return None
2824-
self._host_stager_init_attempted = True
28252819
desc_lens = np.array(
28262820
[int(entry[1]) for entry in self.src_blocks_data], dtype=np.int64
28272821
)
28282822
host_addrs = np.array(
28292823
[int(entry[0]) for entry in self.src_blocks_data], dtype=np.uint64
28302824
)
2825+
return self._maybe_init_host_stager_for_buffers(
2826+
remote_host,
2827+
desc_lens,
2828+
host_addrs,
2829+
list(self.host_xfer_buffers.values()),
2830+
)
2831+
2832+
def _maybe_init_host_stager_for_buffers(
2833+
self,
2834+
remote_host: str,
2835+
desc_lens: np.ndarray,
2836+
host_addrs: np.ndarray,
2837+
host_buffers: list[torch.Tensor],
2838+
) -> HostWriteStager | None:
2839+
"""Enable device staging for same-host reads into host buffers."""
2840+
if remote_host != envs.VLLM_NIXL_SIDE_CHANNEL_HOST:
2841+
return None
2842+
if self._host_stager is not None or self._host_stager_init_attempted:
2843+
return self._host_stager
2844+
stage_bytes = envs.VLLM_NIXL_HOST_STAGE_BYTES
2845+
if stage_bytes <= 0:
2846+
return None
2847+
self._host_stager_init_attempted = True
28312848
try:
28322849
self._host_stager = HostWriteStager(
28332850
desc_lens=desc_lens,
28342851
host_addrs=host_addrs,
2835-
host_buffers=list(self.host_xfer_buffers.values()),
2852+
host_buffers=host_buffers,
28362853
device=torch.device(f"cuda:{self.device_id}"),
28372854
nixl_wrapper=self.nixl_wrapper,
28382855
memory_type=self.nixl_memory_type,
@@ -2853,8 +2870,7 @@ def _get_finished_host_staging(self) -> set[str]:
28532870
if self._host_stager is None:
28542871
return set()
28552872
done, failed = self._host_stager.get_finished()
2856-
for req_id in done:
2857-
self._send_pending_recv_notifs(req_id)
2873+
done = {req_id for req_id in done if self._finish_recv_component(req_id)}
28582874
for req_id in failed:
28592875
self._log_failure(
28602876
failure_type="transfer_failed",
@@ -2863,8 +2879,38 @@ def _get_finished_host_staging(self) -> set[str]:
28632879
)
28642880
self._pending_recv_notifs.pop(req_id, None)
28652881
self._handle_failed_transfer(req_id, None)
2882+
with self._failed_recv_lock:
2883+
drained_failures = {
2884+
req_id
2885+
for req_id in self._failed_recv_pending
2886+
if not self._recving_transfers.get(req_id)
2887+
and not self._host_staging_active(req_id)
2888+
}
2889+
for req_id in drained_failures:
2890+
self._finish_recv_component(req_id)
28662891
return done
28672892

2893+
def _host_staging_active(self, req_id: str) -> bool:
2894+
return (
2895+
self._host_stager is not None and req_id in self._host_stager.active_req_ids
2896+
)
2897+
2898+
def _finish_recv_component(self, req_id: str) -> bool:
2899+
"""Complete a receive once its direct and staged parts are terminal."""
2900+
if self._recving_transfers.get(req_id) or self._host_staging_active(req_id):
2901+
return False
2902+
with self._failed_recv_lock:
2903+
failed = req_id in self._failed_recv_pending
2904+
if failed:
2905+
self._failed_recv_pending.discard(req_id)
2906+
if failed:
2907+
self._report_failed_recv(req_id)
2908+
return False
2909+
if req_id not in self._recving_metadata:
2910+
return False
2911+
self._send_pending_recv_notifs(req_id)
2912+
return True
2913+
28682914
def _pop_done_transfers(
28692915
self, transfers: dict[str, list[int]], *, is_recv: bool
28702916
) -> set[str]:
@@ -2922,18 +2968,9 @@ def _pop_done_transfers(
29222968
transfers[req_id] = in_progress
29232969
continue
29242970
del transfers[req_id]
2925-
if is_recv:
2926-
with self._failed_recv_lock:
2927-
failed = req_id in self._failed_recv_pending
2928-
self._failed_recv_pending.discard(req_id)
2929-
if failed:
2930-
self._report_failed_recv(req_id)
2931-
continue
2932-
if req_id not in self._recving_metadata:
2933-
continue
2971+
if is_recv and not self._finish_recv_component(req_id):
2972+
continue
29342973
done_req_ids.add(req_id)
2935-
if is_recv:
2936-
self._send_pending_recv_notifs(req_id)
29372974
return done_req_ids
29382975

29392976
def _send_pending_recv_notifs(self, req_id: str) -> None:
@@ -2960,7 +2997,7 @@ def _handle_failed_transfer(self, req_id: str, handle: int | None):
29602997
self.xfer_stats.record_failed_transfer()
29612998
if self._host_stager is not None:
29622999
self._host_stager.abort(req_id)
2963-
if req_id in self._host_stager.active_req_ids:
3000+
if self._host_staging_active(req_id):
29643001
with self._failed_recv_lock:
29653002
self._failed_recv_pending.add(req_id)
29663003
return

0 commit comments

Comments
 (0)