Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 76 additions & 18 deletions collectivex/bench/kv_nixl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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]]:
Expand All @@ -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"
Expand All @@ -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))]
Expand All @@ -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

Expand All @@ -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:
Expand All @@ -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)
Expand Down
9 changes: 0 additions & 9 deletions collectivex/bench/kv_workload.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
9 changes: 7 additions & 2 deletions collectivex/tests/test_kv_workload.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand Down