From c64e2210574e9e5844978a2118420a26431c09a1 Mon Sep 17 00:00:00 2001 From: Ilia Yastrebov Date: Mon, 28 Sep 2026 15:58:27 +0200 Subject: [PATCH] CollectiveX: use NIXL's two-step prepped-list API for kv-transfer --- collectivex/bench/kv_nixl.py | 94 ++++++++++++++++++++++----- collectivex/bench/kv_workload.py | 9 --- collectivex/tests/test_kv_workload.py | 9 ++- 3 files changed, 83 insertions(+), 29 deletions(-) diff --git a/collectivex/bench/kv_nixl.py b/collectivex/bench/kv_nixl.py index 07148ca987..b81c9f4df0 100644 --- a/collectivex/bench/kv_nixl.py +++ b/collectivex/bench/kv_nixl.py @@ -7,7 +7,9 @@ adapter needs no port and no listener race. Remote descriptors are built locally from the peer's published pool base; both block tables are seed-keyed, the same information a decode worker gets from the -prefill side's block table message. +prefill side's block table message. Transfers use NIXL's two-step API: each +side's pool descriptor list is prepped once, and every request is a +`make_prepped_xfer` over its block tables as row indices into those lists. """ from __future__ import annotations @@ -17,7 +19,6 @@ import numpy as np -import kv_workload from kv_backend import KVBackend, library_version # b300's CX NICs refuse cuda registrations somewhere between 7083 and 8847 MiB @@ -27,6 +28,7 @@ # its own packed-block grid so no transfer descriptor straddles two pieces. REG_CHUNK_BYTES = 4 << 30 +LOCAL_SIDE = "NIXL_INIT_AGENT" def reg_spans(nbytes: int, layout, cap: int = REG_CHUNK_BYTES) -> list[tuple[int, int]]: @@ -50,6 +52,32 @@ def reg_spans(nbytes: int, layout, return spans +def pool_desc_array(base: int, layout, dev: int) -> np.ndarray: + """(addr, len, devId) uint64 rows, one per packed block of ``layout``, in layout order.""" + rows = [] + for region_base, packed, region_nbytes in layout: + blocks = region_nbytes // packed + region = np.empty((blocks, 3), dtype=np.uint64) + region[:, 0] = (np.uint64(base + region_base) + + np.arange(blocks, dtype=np.uint64) * np.uint64(packed)) + region[:, 1] = packed + region[:, 2] = dev + rows.append(region) + return np.concatenate(rows) + + +def desc_indices(layout, cfg: dict, tables: dict) -> np.ndarray: + """Rows of one request's blocks in the pool list (uint32, C-contiguous).""" + first_row = {} + row = 0 + for region_base, packed, region_nbytes in layout: + first_row[region_base] = row + row += region_nbytes // packed + return np.ascontiguousarray(np.concatenate([ + np.asarray(tables[region["name"]], dtype=np.int64) + first_row[region["base"]] + for region in cfg["regions"]]), dtype=np.uint32) + + class NIXLBackend(KVBackend): name = "nixl" maturity = "production" @@ -67,12 +95,16 @@ def __init__(self, args, role, device): self._agent = nixl_agent(role, nixl_agent_config(True, False, 0, backends=[self.transport])) self._handles = [] + self._dlists = [] self._pool = None self._bulk = None self._peer = None + self._layout = None + self._pool_lists = None + self._bulk_lists = {} def register(self, pool, bulk, reg_layout=None) -> None: - self._pool, self._bulk = pool, bulk + self._pool, self._bulk, self._layout = pool, bulk, reg_layout entries = [(pool.ptr + off, length, pool.device, f"pool{i}") for i, (off, length) in enumerate(reg_spans(pool.nbytes, reg_layout))] @@ -96,15 +128,36 @@ def connect(self, peer: dict) -> None: remote = self._agent.add_remote_agent(peer["agent"]) self._remote_name = remote.decode() if isinstance(remote, (bytes, bytearray)) else str(remote) - def _make(self, local_np: np.ndarray, remote_np: np.ndarray, op: str): - start = time.perf_counter() - local_descs = self._agent.get_xfer_descs(local_np, mem_type="cuda") - remote_descs = self._agent.get_xfer_descs(remote_np, mem_type="cuda") - handle = self._agent.initialize_xfer( + def _prep(self, side: str, descs: np.ndarray): + handle = self._agent.prep_xfer_dlist(side, descs, mem_type="cuda") + self._dlists.append(handle) + return handle + + def _pool_sides(self): + if self._pool_lists is None: + local = pool_desc_array(self._pool.ptr, self._layout, self._pool.device) + remote = pool_desc_array(self._peer["pool_base"], self._layout, self._peer["dev"]) + self._pool_lists = (self._prep(LOCAL_SIDE, local), + self._prep(self._remote_name, remote)) + return self._pool_lists + + def _bulk_sides(self, nbytes: int): + if nbytes not in self._bulk_lists: + local = np.array([[self._bulk.ptr, nbytes, self._bulk.device]], dtype=np.uint64) + remote = np.array([[self._peer["bulk_base"], nbytes, self._peer["dev"]]], + dtype=np.uint64) + self._bulk_lists[nbytes] = (self._prep(LOCAL_SIDE, local), + self._prep(self._remote_name, remote)) + return self._bulk_lists[nbytes] + + def _make(self, sides, local_idx: np.ndarray, remote_idx: np.ndarray, op: str, + started: float): + local_side, remote_side = sides + handle = self._agent.make_prepped_xfer( "READ" if op == "pull" else "WRITE", - local_descs, remote_descs, self._remote_name, + local_side, local_idx, remote_side, remote_idx, ) - prep_s = time.perf_counter() - start + prep_s = time.perf_counter() - started self._handles.append(handle) agent = self._agent @@ -123,16 +176,16 @@ def wait(): return post, wait, prep_s def make_paged(self, cfg, op, local_tables, remote_tables): - local_np = kv_workload.desc_array(self._pool.ptr, cfg, local_tables, self._pool.device) - remote_np = kv_workload.desc_array(self._peer["pool_base"], cfg, remote_tables, - self._peer["dev"]) - return self._make(local_np, remote_np, op) + sides = self._pool_sides() + start = time.perf_counter() + return self._make(sides, desc_indices(self._layout, cfg, local_tables), + desc_indices(self._layout, cfg, remote_tables), op, start) def make_bulk(self, nbytes, op): - local_np = np.array([[self._bulk.ptr, nbytes, self._bulk.device]], dtype=np.uint64) - remote_np = np.array([[self._peer["bulk_base"], nbytes, self._peer["dev"]]], - dtype=np.uint64) - return self._make(local_np, remote_np, op) + sides = self._bulk_sides(nbytes) + start = time.perf_counter() + first = np.zeros(1, dtype=np.uint32) + return self._make(sides, first, first, op, start) def release(self) -> None: for handle in self._handles: @@ -144,6 +197,11 @@ def release(self) -> None: def teardown(self) -> None: self.release() + for handle in self._dlists: + try: + self._agent.release_dlist_handle(handle) + except Exception: + pass if self._peer is not None: try: self._agent.remove_remote_agent(self._remote_name) diff --git a/collectivex/bench/kv_workload.py b/collectivex/bench/kv_workload.py index 755a63b24c..6c2d227971 100644 --- a/collectivex/bench/kv_workload.py +++ b/collectivex/bench/kv_workload.py @@ -157,15 +157,6 @@ def desc_sizes(cfg: dict) -> np.ndarray: ]) -def desc_array(base: int, cfg: dict, tables: dict, dev: int) -> np.ndarray: - """(addr, len, devId) uint64 rows for descriptor-list APIs (NIXL's numpy form).""" - out = np.empty((cfg["descs"], 3), dtype=np.uint64) - out[:, 0] = np.uint64(base) + page_offsets(cfg, tables) - out[:, 1] = desc_sizes(cfg) - out[:, 2] = dev - return out - - # The pattern depends on (offset >> 8) mod 256 only, so it repeats every 64 KiB. PATTERN_PERIOD = 1 << 16 diff --git a/collectivex/tests/test_kv_workload.py b/collectivex/tests/test_kv_workload.py index 17a39d7a85..12b75bba7e 100644 --- a/collectivex/tests/test_kv_workload.py +++ b/collectivex/tests/test_kv_workload.py @@ -131,13 +131,18 @@ def test_a_request_beyond_the_pool_fails_closed(self): with self.assertRaises(ValueError): kv_workload.block_table(cfg, 1, request=8) - def test_desc_array_carries_per_region_packed_sizes(self): + def test_prepped_rows_carry_per_region_packed_sizes(self): + import kv_nixl + cfg = dict(regions=[ dict(name="a", packed_bytes=256, blocks_req=2, pool_blocks=4, base=0), dict(name="b", packed_bytes=132, blocks_req=1, pool_blocks=4, base=1024), ], descs=3) + layout = [(0, 256, 4 * 256), (1024, 132, 4 * 132)] tables = {"a": np.array([1, 3]), "b": np.array([2])} - descs = kv_workload.desc_array(10_000, cfg, tables, dev=5) + indices = kv_nixl.desc_indices(layout, cfg, tables) + self.assertEqual(indices.tolist(), [1, 3, 4 + 2]) + descs = kv_nixl.pool_desc_array(10_000, layout, dev=5)[indices] self.assertEqual(descs[:, 0].tolist(), [10_000 + 256, 10_000 + 768, 10_000 + 1024 + 264]) self.assertEqual(descs[:, 1].tolist(), [256, 256, 132])