From 2a2236507384c5215e1761e0ec5462cbcdf4bb6e Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:26:42 +0800 Subject: [PATCH 01/21] transfer: shm submit/result rings and a node-local shm TE process; per-process op/graph id ranges Every DP rank's cache engine talks to ONE TransferEngine subprocess per node over shared memory instead of each rank owning a TE: * flexkv/transfer/shm_channel.py: fixed-slot submit ring (fragmented pickles) and fixed-width CompletedOp result ring per client, plus a futex control block the TE parks on when idle. * flexkv/transfer/shm_channel_handle.py: TransferManagerShmChannelHandle (client side) and te_shm_main, which brings up the TransferManager, polls all N submit rings, forwards graphs and routes completions back by client. * flexkv/transfer_manager.py: TransferManagerShmTEProcess owns that subprocess; CUDA_VISIBLE_DEVICES is cleared for it only when the deployment spans more than one GPU. * flexkv/common/transfer.py: TransferOp.set_op_id_range and TransferOpGraph.set_graph_id_range give each CE process a disjoint id range so submissions to the shared TE never collide. * flexkv/transfer/transfer_engine.py: no FlexKV peer transfer worker in radixshmem mode; the radix-server pulls peer blocks itself. Reconciled with main during the rebase: * main replaced the id locks with itertools.count(); the ranges are now ``start + counter % size`` on top of that counter, still lock-free. * CompletedOp gained wait_ms/xfer_ms/e2e_ms on main (#297); the result ring does not carry them, so shm-channel completions report 0.0. * KVServer and TransferManagerOnRemote follow main #287 and inherit CUDA_VISIBLE_DEVICES: server-client mode and MPS are still in use, and neither launcher is on the radixshmem path (that path spawns TransferManagerShmTEProcess, which keeps its own total_gpus > 1 rule). The branch's 17ccce7 revert of #287 (radixshmem_rebase_v1.bak-20260921) is not carried over. Rebased from radixshmem_rebase_v1 (68 commits, kept as radixshmem_rebase_v1.bak-20260921) onto main 738ddc1 as one squash-merge, then split by subsystem; this is part 1/7. Co-authored-by: linhu-nv Co-authored-by: Iris Ge Co-authored-by: Hao Xu --- flexkv/common/transfer.py | 44 ++- flexkv/transfer/shm_channel.py | 477 ++++++++++++++++++++++++++ flexkv/transfer/shm_channel_handle.py | 350 +++++++++++++++++++ flexkv/transfer/transfer_engine.py | 9 +- flexkv/transfer_manager.py | 114 +++++- tests/test_shm_channel.py | 259 ++++++++++++++ 6 files changed, 1246 insertions(+), 7 deletions(-) create mode 100644 flexkv/transfer/shm_channel.py create mode 100644 flexkv/transfer/shm_channel_handle.py create mode 100644 tests/test_shm_channel.py diff --git a/flexkv/common/transfer.py b/flexkv/common/transfer.py index 9d43cdfad..89b824d3b 100644 --- a/flexkv/common/transfer.py +++ b/flexkv/common/transfer.py @@ -172,6 +172,21 @@ class TransferOp: # the cost. Kept as a ClassVar so ids stay global across all graphs, which # merge_to_batch_graph relies on when it mixes ops from many tasks. _op_id_counter: ClassVar["itertools.count"] = itertools.count() + # Per-process disjoint range, set by `set_op_id_range()`. Default is the + # full int64 positive range, preserving single-CE behavior. The radix-shmem + # multi-DP path partitions this so 8 CE procs sharing one TE never collide. + # An id is ``start + counter % size``: still lock-free, and it wraps inside + # the range instead of running into a neighbour's. + _op_id_range_start: ClassVar[int] = 0 + _op_id_range_size: ClassVar[int] = 1 << 62 + + @classmethod + def set_op_id_range(cls, start: int, end: int) -> None: + """Restrict generated op_ids to [start, end). Call before any op is + created in this process; it restarts the counter.""" + cls._op_id_range_start = start + cls._op_id_range_size = end - start + cls._op_id_counter = itertools.count() op_id: int = field(init=False) graph_id: int @@ -259,7 +274,8 @@ def __post_init__(self, is_swa: Optional[bool] = None) -> None: raise ValueError(f"src_block_ids and dst_block_ids must have the same number of physical blocks, but got " f"src_block_ids.size={src.size}, " f"dst_block_ids.size={dst.size}") - self.op_id = next(TransferOp._op_id_counter) + self.op_id = (TransferOp._op_id_range_start + + next(TransferOp._op_id_counter) % TransferOp._op_id_range_size) assert src.dtype == _INT64 assert dst.dtype == _INT64 self.valid_block_num = src.size @@ -329,9 +345,18 @@ class TransferOpGraph: # Lock-free for the same reason as TransferOp._op_id_counter: a C-level # __next__ that cannot be interrupted mid-increment. _graph_id_counter = itertools.count() + # Per-process DP-aware range, set via set_graph_id_range(start, end). Used + # so multiple CE processes that share a single TE don't collide on + # graph_id. Default range is the original (0, 2**62), preserving behavior + # in the single-CE path. Same ``start + counter % size`` scheme as + # TransferOp. + _graph_id_range_start = 0 + _graph_id_range_size = 1 << 62 def __init__(self) -> None: - self.graph_id = next(TransferOpGraph._graph_id_counter) + self.graph_id = (TransferOpGraph._graph_id_range_start + + next(TransferOpGraph._graph_id_counter) + % TransferOpGraph._graph_id_range_size) self._op_map: Dict[int, TransferOp] = {} self._ready_ops: Set[int] = set() self._trigger_ops: Set[int] = set() @@ -346,7 +371,18 @@ def __init__(self) -> None: @classmethod def _get_graph_id(cls) -> int: """Kept as the named entry point; __init__ inlines the same counter.""" - return next(cls._graph_id_counter) + return (cls._graph_id_range_start + + next(cls._graph_id_counter) % cls._graph_id_range_size) + + @classmethod + def set_graph_id_range(cls, start: int, end: int) -> None: + """Restrict generated graph_ids to [start, end). Used by the + radix-shmem multi-DP path to give each CE process a disjoint range. + Call before any graph is created in this process; it restarts the + counter.""" + cls._graph_id_range_start = start + cls._graph_id_range_size = end - start + cls._graph_id_counter = itertools.count() def set_graph_id(self, graph_id: int) -> None: self.graph_id = graph_id @@ -1241,7 +1277,7 @@ def merge_to_batch_graph(batch_id: int, put_sinks.append(merged_swa_d2h_op.op_id) if not put_sinks: # No D2H sink: wait for every independent full-KV / SWA leaf - # (H2DISK and/or H2REMOTE). + # (H2DISK and/or H2REMOTE). for op in (merged_h2disk_op, merged_swa_h2disk_op, merged_h2remote_op, merged_swa_h2remote_op): if op is not None: diff --git a/flexkv/transfer/shm_channel.py b/flexkv/transfer/shm_channel.py new file mode 100644 index 000000000..652a2defa --- /dev/null +++ b/flexkv/transfer/shm_channel.py @@ -0,0 +1,477 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +Shared-memory IPC channel for CacheEngine ↔ TransferEngine communication in +multi-DP FlexKV. + +Adapted from PR #144 (commit 5e262ca, originally for DPClient ↔ KVServer). +Slimmed down for the CE↔TE use case: + - Per-channel size is small (256 KB ring + 256 KB sync) because the only + payloads are pickled TransferOpGraph (submit) and CompletedOp lists (wait). + - Sync request/response slot is dropped — CE↔TE is fully fire-and-forget on + both directions; submit is async, completions are pushed asynchronously. + - Two SPSC ring buffers per channel: `submit` (CE→TE) and `result` (TE→CE). + Each side futex-waits on its own counter when the ring is empty. + +Layout per channel (one /dev/shm file per CE): + [0..64) submit_write_pos (uint64, CE writes) + [64..128) submit_read_pos (uint64, TE writes) + [128..192) submit_wake (int32, CE bumps + futex_wake) + [192..256) result_write_pos (uint64, TE writes) + [256..320) result_read_pos (uint64, CE writes) + [320..384) result_wake (int32, TE bumps + futex_wake) + [384..ring_off) reserved + [ring_off ..) submit ring (slot * SUBMIT_SLOTS) + result ring (slot * RESULT_SLOTS) + +A single `ShmControlBlock` (separate /dev/shm file) carries a global wake +counter that the TE polls when it has more than one channel attached, so it +can sleep idly without per-channel futex wait. +""" +from __future__ import annotations + +import ctypes +import ctypes.util +import mmap +import os +import pickle +import platform +import struct +from typing import Any, List, Optional + +# ── Linux futex wrappers ──────────────────────────────────────────────── + +_libc = ctypes.CDLL(ctypes.util.find_library("c"), use_errno=True) + +# futex syscall number is arch-specific; pick at import time. +_MACHINE = platform.machine().lower() +if _MACHINE in ("x86_64", "amd64"): + _SYS_FUTEX = 202 +elif _MACHINE in ("aarch64", "arm64"): + _SYS_FUTEX = 98 +else: # pragma: no cover + raise RuntimeError(f"Unsupported machine for futex syscall: {_MACHINE}") +_FUTEX_WAIT = 0 +_FUTEX_WAKE = 1 + + +def _futex_wait(addr: int, expected: int, timeout_ns: Optional[int] = None) -> int: + if timeout_ns is None: + return _libc.syscall( + _SYS_FUTEX, ctypes.c_void_p(addr), + _FUTEX_WAIT, ctypes.c_int(expected), + ctypes.c_void_p(0), ctypes.c_void_p(0), ctypes.c_int(0), + ) + # struct timespec + ts = (ctypes.c_long * 2)(timeout_ns // 1_000_000_000, + timeout_ns % 1_000_000_000) + return _libc.syscall( + _SYS_FUTEX, ctypes.c_void_p(addr), + _FUTEX_WAIT, ctypes.c_int(expected), + ctypes.byref(ts), ctypes.c_void_p(0), ctypes.c_int(0), + ) + + +def _futex_wake(addr: int, count: int = 1) -> int: + return _libc.syscall( + _SYS_FUTEX, ctypes.c_void_p(addr), + _FUTEX_WAKE, ctypes.c_int(count), + ctypes.c_void_p(0), ctypes.c_void_p(0), ctypes.c_int(0), + ) + + +# ── Layout constants ──────────────────────────────────────────────────── + +_CL = 64 # cache line + +# Header lives in the first 6 cache lines; ring data starts on a page boundary. +OFF_SUBMIT_W = 0 * _CL +OFF_SUBMIT_R = 1 * _CL +OFF_SUBMIT_WAKE = 2 * _CL +OFF_RESULT_W = 3 * _CL +OFF_RESULT_R = 4 * _CL +OFF_RESULT_WAKE = 5 * _CL +HEADER_SIZE = 6 * _CL # 384 B + +# Submit ring holds pickled TransferOpGraphs, fragmented across slots so a payload +# larger than one slot spans several (32 KB × 8192 = 256 MB, ~256 MB max message). +# Fragmentation retired the old "must fit one slot" constraint that made high-QPS +# as_batch=True graphs overflow (observed 152 KB @2500 with prefetch). +DEFAULT_SUBMIT_SLOTS = 8192 # power of 2 +DEFAULT_SUBMIT_SLOT_SIZE = 32 * 1024 +DEFAULT_SLOT_SIZE = DEFAULT_SUBMIT_SLOT_SIZE # back-compat alias for `slot_size=` + +# Result ring holds one fixed-width CompletedOp record per slot (64 B × 65536 = 4 MB). +DEFAULT_RESULT_SLOTS = 65536 # power of 2 +DEFAULT_RESULT_SLOT_SIZE = 64 # cache line; one 29 B CompletedOp record + +_PAGE = 4096 + + +def _round_up(x: int, m: int) -> int: + return (x + m - 1) // m * m + + +# Submit-ring fragment header (per slot): payload bytes in this slot (u32) + a +# last-fragment flag (u8). A message is the concatenation of fragments up to and +# including the one with is_last=1. +_FRAG_HDR = struct.Struct(" bytes: + """Pack a CompletedOp into its fixed-width record.""" + tt = op.transfer_type + tt_idx = _TT_NONE if tt is None else _TT_NAME_TO_IDX[tt] + flags = 1 if getattr(op, "failed", False) else 0 + return _COMPLETED_OP.pack( + op.graph_id, op.op_id, tt_idx, op.num_blocks, op.num_bytes, flags, + ) + + +def decode_completed_op(buf: Any, off: int) -> Any: + """Unpack a CompletedOp record from `buf` at byte offset `off`.""" + from flexkv.common.transfer import CompletedOp + graph_id, op_id, tt_idx, num_blocks, num_bytes, flags = \ + _COMPLETED_OP.unpack_from(buf, off) + tt = None if tt_idx == _TT_NONE else _TT_NAMES[tt_idx] + return CompletedOp( + graph_id=graph_id, + op_id=op_id, + transfer_type=tt, + num_blocks=num_blocks, + num_bytes=num_bytes, + failed=bool(flags & 1), + ) + + +# ── ShmControlBlock ───────────────────────────────────────────────────── + +CTRL_WAKE = 0 +CTRL_READY = _CL +CTRL_SIZE = _PAGE + + +def _safe_id(server_id: str) -> str: + return server_id.replace("/", "_").replace(":", "_").strip("_") + + +class ShmControlBlock: + """Optional global wake counter used by the TE when polling N channels.""" + + def __init__(self, server_id: str, create: bool = False): + self.server_id = server_id + safe = _safe_id(server_id) + self.shm_path = f"/dev/shm/flexkv_te_ctrl_{safe}" + + if create: + fd = os.open(self.shm_path, os.O_CREAT | os.O_RDWR, 0o666) + os.ftruncate(fd, CTRL_SIZE) + self.buf = mmap.mmap(fd, CTRL_SIZE) + os.close(fd) + self.buf[:] = b"\x00" * CTRL_SIZE + else: + fd = os.open(self.shm_path, os.O_RDWR) + self.buf = mmap.mmap(fd, CTRL_SIZE) + os.close(fd) + + self._base = ctypes.addressof(ctypes.c_char.from_buffer(self.buf)) + self._wake = ctypes.c_int32.from_address(self._base + CTRL_WAKE) + self._ready = ctypes.c_int32.from_address(self._base + CTRL_READY) + + @property + def _wake_addr(self) -> int: + return self._base + CTRL_WAKE + + @property + def _ready_addr(self) -> int: + return self._base + CTRL_READY + + def notify(self) -> None: + # Read-modify-write is not atomic across processes; the TE uses snapshot + # comparison so any change wakes it. + self._wake.value += 1 + _futex_wake(self._wake_addr, 1) + + def get_wake(self) -> int: + return self._wake.value + + def wait(self, expected: int, timeout_ns: Optional[int] = None) -> None: + _futex_wait(self._wake_addr, expected, timeout_ns) + + def set_ready(self) -> None: + self._ready.value = 1 + _futex_wake(self._ready_addr, 0x7FFFFFFF) + + def wait_ready(self, timeout_s: float = 60.0) -> bool: + import time + deadline = time.monotonic() + timeout_s + while time.monotonic() < deadline: + if self._ready.value != 0: + return True + _futex_wait(self._ready_addr, 0, + timeout_ns=int(0.5 * 1_000_000_000)) + return False + + def close(self) -> None: + if self.buf is not None: + self.buf.close() + self.buf = None + + def unlink(self) -> None: + try: + os.unlink(self.shm_path) + except FileNotFoundError: + pass + + +# ── ShmChannel (CE ↔ TE) ──────────────────────────────────────────────── + +class ShmChannel: + """One bi-directional channel between a single CE and the TE. + + Two SPSC rings: CE→TE submit, TE→CE result. Each side futex-waits on its + own wake counter when its consumer ring is empty. + """ + + def __init__(self, + server_id: str, + channel_id: int, + create: bool = False, + submit_slots: int = DEFAULT_SUBMIT_SLOTS, + result_slots: int = DEFAULT_RESULT_SLOTS, + slot_size: int = DEFAULT_SUBMIT_SLOT_SIZE, + result_slot_size: int = DEFAULT_RESULT_SLOT_SIZE): + assert submit_slots & (submit_slots - 1) == 0, \ + "submit_slots must be power of 2" + assert result_slots & (result_slots - 1) == 0, \ + "result_slots must be power of 2" + assert result_slot_size >= COMPLETED_OP_WIRE_SIZE, \ + f"result_slot_size {result_slot_size} < CompletedOp record " \ + f"{COMPLETED_OP_WIRE_SIZE}" + + self.channel_id = channel_id + self.submit_slots = submit_slots + self.result_slots = result_slots + self.slot_size = slot_size # submit ring; result ring uses result_slot_size + self.result_slot_size = result_slot_size + + safe = _safe_id(server_id) + self.shm_path = f"/dev/shm/flexkv_te_ch_{safe}_{channel_id}" + + # Lay out: header -> aligned to page -> submit ring -> result ring. + self._submit_off = _round_up(HEADER_SIZE, _PAGE) + self._result_off = self._submit_off + submit_slots * slot_size + total = self._result_off + result_slots * result_slot_size + + self.total_size = total + + if create: + fd = os.open(self.shm_path, os.O_CREAT | os.O_RDWR, 0o666) + os.ftruncate(fd, total) + self.buf = mmap.mmap(fd, total) + os.close(fd) + self.buf[:HEADER_SIZE] = b"\x00" * HEADER_SIZE + else: + fd = os.open(self.shm_path, os.O_RDWR) + self.buf = mmap.mmap(fd, total) + os.close(fd) + + self._base = ctypes.addressof(ctypes.c_char.from_buffer(self.buf)) + self._submit_w = ctypes.c_uint64.from_address(self._base + OFF_SUBMIT_W) + self._submit_r = ctypes.c_uint64.from_address(self._base + OFF_SUBMIT_R) + self._submit_wake = ctypes.c_int32.from_address(self._base + OFF_SUBMIT_WAKE) + self._result_w = ctypes.c_uint64.from_address(self._base + OFF_RESULT_W) + self._result_r = ctypes.c_uint64.from_address(self._base + OFF_RESULT_R) + self._result_wake = ctypes.c_int32.from_address(self._base + OFF_RESULT_WAKE) + + # ---- futex helpers ---- + + @property + def _submit_wake_addr(self) -> int: + return self._base + OFF_SUBMIT_WAKE + + @property + def _result_wake_addr(self) -> int: + return self._base + OFF_RESULT_WAKE + + def _bump_wake(self, ptr: ctypes.c_int32, addr: int) -> None: + ptr.value += 1 + _futex_wake(addr, 1) + + # ---- ring helpers ---- + + def _ring_full(self, w: int, r: int, slots: int) -> bool: + return ((w + 1) & (slots - 1)) == r + + def _ring_used(self, w: int, r: int, slots: int) -> int: + return (w - r) & (slots - 1) + + # ---- CE side: submit + recv result ---- + + def submit_send(self, payload: Any) -> None: + """Enqueue a payload to TE, fragmenting it across slots if it exceeds one. + + All fragments are written first, then submit_w is advanced once, so the TE + never observes a partial message. Spins+yields until enough contiguous + slots are free.""" + blob = pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL) + slots = self.submit_slots + body = self.slot_size - _FRAG_HDR_SIZE + nfrag = max(1, (len(blob) + body - 1) // body) + if nfrag > slots - 1: + raise ValueError( + f"payload needs {nfrag} fragments but submit ring holds " + f"{slots - 1}; raise submit_slots or slot_size") + + wp = self._submit_w.value + spin = 0 + while self._ring_used(wp, self._submit_r.value, slots) + nfrag > slots - 1: + if spin > 1_000_000: + raise RuntimeError("shm channel submit ring full") + if spin > 1000: + os.sched_yield() + spin += 1 + + w = wp + for i in range(nfrag): + chunk = blob[i * body:(i + 1) * body] + off = self._submit_off + w * self.slot_size + self.buf[off:off + _FRAG_HDR_SIZE] = _FRAG_HDR.pack( + len(chunk), 1 if i == nfrag - 1 else 0) + self.buf[off + _FRAG_HDR_SIZE:off + _FRAG_HDR_SIZE + len(chunk)] = chunk + w = (w + 1) & (slots - 1) + self._submit_w.value = w + self._bump_wake(self._submit_wake, self._submit_wake_addr) + + def result_recv(self, timeout_s: Optional[float] = None) -> List[Any]: + """Drain pending TE→CE results, one CompletedOp per slot. Blocks up to + `timeout_s` if empty.""" + out: List[Any] = [] + rp = self._result_r.value + wp = self._result_w.value + slots = self.result_slots + if rp == wp and timeout_s is not None and timeout_s > 0: + wake = self._result_wake.value + # Re-check; TE might have arrived between read and wait. + wp = self._result_w.value + if rp == wp: + if timeout_s == float("inf"): + _futex_wait(self._result_wake_addr, wake) + else: + _futex_wait(self._result_wake_addr, wake, + timeout_ns=int(timeout_s * 1_000_000_000)) + wp = self._result_w.value + + while rp != wp: + out.append(decode_completed_op( + self.buf, self._result_off + rp * self.result_slot_size)) + rp = (rp + 1) & (slots - 1) + if out: + self._result_r.value = rp + return out + + # ---- TE side: recv submit + send result ---- + + def submit_recv(self) -> List[Any]: + """Drain CE→TE submissions (non-blocking), reassembling fragmented + messages. submit_r is advanced only past fully-received messages.""" + out: List[Any] = [] + rp = self._submit_r.value + wp = self._submit_w.value + slots = self.submit_slots + frags: List[bytes] = [] + while rp != wp: + off = self._submit_off + rp * self.slot_size + n, is_last = _FRAG_HDR.unpack_from(self.buf, off) + frags.append(bytes(self.buf[off + _FRAG_HDR_SIZE: + off + _FRAG_HDR_SIZE + n])) + rp = (rp + 1) & (slots - 1) + if is_last: + out.append(pickle.loads(b"".join(frags))) + frags.clear() + self._submit_r.value = rp # release this message's slots + return out + + def result_send(self, ops: List[Any]) -> None: + """Enqueue a batch of CompletedOps, one fixed-width record per slot, and + wake the CE once at the end. Spins if the ring fills rather than dropping a + completion (which would hang the owning task); with 65536 slots that is + effectively unreachable.""" + if not ops: + return + slots = self.result_slots + slot_sz = self.result_slot_size + base = self._result_off + wp = self._result_w.value + warned = False + for op in ops: + if self._ring_full(wp, self._result_r.value, slots): + # Full mid-batch: publish+wake so the CE drains, then spin. + self._bump_wake(self._result_wake, self._result_wake_addr) + spin = 0 + while self._ring_full(wp, self._result_r.value, slots): + if spin > 1000: + os.sched_yield() + spin += 1 + if spin % 5_000_000 == 0 and not warned: + try: + from flexkv.common.debug import flexkv_logger + flexkv_logger.error( + f"shm channel result ring stuck full " + f"(slots={slots}); is the CE consumer alive?" + ) + except Exception: + pass + warned = True + off = base + wp * slot_sz + self.buf[off:off + COMPLETED_OP_WIRE_SIZE] = encode_completed_op(op) + wp = (wp + 1) & (slots - 1) + self._result_w.value = wp + self._bump_wake(self._result_wake, self._result_wake_addr) + + # ---- Submit-wake fileno: lets TE selector wait on this channel ---- + + @property + def submit_wake_fd(self) -> int: + # We don't have a real eventfd; selector users should poll get_wake() + # delta + futex_wait via ShmControlBlock. Returning -1 signals "no fd". + return -1 + + # ---- Lifecycle ---- + + def close(self) -> None: + if self.buf is not None: + self.buf.close() + self.buf = None + + def unlink(self) -> None: + try: + os.unlink(self.shm_path) + except FileNotFoundError: + pass + + def __del__(self) -> None: + try: + self.close() + except Exception: + pass diff --git a/flexkv/transfer/shm_channel_handle.py b/flexkv/transfer/shm_channel_handle.py new file mode 100644 index 000000000..4b514475a --- /dev/null +++ b/flexkv/transfer/shm_channel_handle.py @@ -0,0 +1,350 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +Shared-memory variant of TransferManagerHandle for the multi-DP path. + +Architecture: +- N CE processes each hold one `TransferManagerShmChannelHandle` connected to + the single TE process via a `ShmChannel` named after the (server_id, + channel_id). +- The TE process runs a multi-channel dispatcher loop (`_te_shm_main`) that + polls all N submit rings, hands graphs to the underlying TransferManager, + and routes each completed op back to its originating channel via a + `graph_id → channel_id` map. +""" +from __future__ import annotations + +import os +import queue +import threading +import time +from typing import Dict, List, Optional, Tuple + +import nvtx + +from flexkv.common.config import CacheConfig, ModelConfig +from flexkv.common.debug import flexkv_logger +from flexkv.common.transfer import CompletedOp, TransferOpGraph +from flexkv.transfer.shm_channel import ShmChannel, ShmControlBlock + + +# Wire format: small dicts so we don't pay for full TransferOpGraph re-pickling +# more than necessary. Submit messages carry the graph itself; we'll let pickle +# handle the structure. + +class _SubmitMsg: + __slots__ = ("graph", "task_end_op_id", "is_batch") + + def __init__(self, graph, task_end_op_id: int = -1, is_batch: bool = False): + self.graph = graph + self.task_end_op_id = task_end_op_id + self.is_batch = is_batch + + +# The result ring carries fixed-width CompletedOp records directly (no wrapper). + + +# CE-side handle --------------------------------------------------------- + +class TransferManagerShmChannelHandle: + """CE-side handle that submits transfer graphs to a shared TE via shmem.""" + + def __init__(self, + model_config: ModelConfig, + cache_config: CacheConfig, + server_id: str, + channel_id: int, + file_wait_timeout_s: float = 60.0): + from flexkv.transfer.shm_channel import _safe_id + + self.model_config = model_config + self.cache_config = cache_config + self.server_id = server_id + self.channel_id = channel_id + + safe = _safe_id(server_id) + self._ctrl_path = f"/dev/shm/flexkv_te_ctrl_{safe}" + self._ch_path = f"/dev/shm/flexkv_te_ch_{safe}_{channel_id}" + self._file_wait_timeout_s = file_wait_timeout_s + + # Poll for shm files (created in TE subprocess by setup_channels()). + # Existence does NOT mean the TM is initialized — that's signalled + # by the ctrl ready flag and observed via `is_ready()`. + deadline = time.time() + file_wait_timeout_s + while time.time() < deadline: + if os.path.exists(self._ctrl_path) and os.path.exists(self._ch_path): + break + time.sleep(0.05) + else: + raise RuntimeError( + f"Timed out waiting for TE shm files: " + f"{self._ctrl_path}, {self._ch_path}" + ) + # Attaching is non-blocking once files exist. + self._ctrl = ShmControlBlock(server_id, create=False) + self._channel = ShmChannel(server_id, channel_id, create=False) + + # TransferManagerHandleBase interface ------------------------------- + + def start(self) -> None: + # Nothing to do — TE creates and starts the channel. + pass + + def is_ready(self) -> bool: + # The TE subprocess flips the ctrl ready flag after the + # TransferManager (incl. GPU registration) finishes initializing. + return self._ctrl._ready.value != 0 + + def submit(self, transfer_graph: TransferOpGraph, + task_end_op_id: int = -1) -> None: + nvtx_range = nvtx.start_range( + message="TransferManagerShmChannelHandle.submit", color="green" + ) + self._channel.submit_send(_SubmitMsg(transfer_graph, task_end_op_id)) + self._ctrl.notify() # wake the TE: an idle _poll_submits parks on the ctrl futex for up to 100 ms + nvtx.end_range(nvtx_range) + + def submit_batch(self, transfer_graphs: List[TransferOpGraph]) -> None: + # Send each graph as its own submit message — keeps the TE side + # simple. Could be batched into a list for fewer pickle calls if + # benchmarks show it matters. + for g in transfer_graphs: + self._channel.submit_send(_SubmitMsg(g, -1, is_batch=True)) + if transfer_graphs: + self._ctrl.notify() # see submit() + + def wait(self, timeout: Optional[float] = None) -> List[CompletedOp]: + if timeout is None: + timeout = 0.0 + out: List[CompletedOp] = self._channel.result_recv(timeout_s=timeout) + if out and os.environ.get("FLEXKV_TRACE_TE", "0") == "1": + completed_graphs = sorted({op.graph_id for op in out + if op.is_graph_completed()}) + flexkv_logger.info( + f"[TE-TRACE] CE recv ch={self.channel_id} " + f"completed_graphs={completed_graphs} n_ops={len(out)}" + ) + return out + + def shutdown(self) -> None: + try: + self._channel.close() + except Exception: + pass + try: + self._ctrl.close() + except Exception: + pass + + def __del__(self) -> None: + try: + self.shutdown() + except Exception: + pass + + +# TE-side multi-channel dispatcher -------------------------------------- + +class _TEShmDispatcher: + """Polls N channels, forwards submits to TransferManager, routes results. + + Two-phase startup: + 1) `setup_channels()` creates the control block + per-channel shm files + and sets the ready flag. CE-side handles can attach as soon as this + returns. TM does not need to exist yet. + 2) `start_dispatch(transfer_manager)` launches the polling threads. Call + this after the TM has finished initializing. + """ + + def __init__(self, server_id: str, num_channels: int): + self._tm = None + self._server_id = server_id + self._num_channels = num_channels + self._ctrl: Optional[ShmControlBlock] = None + self._channels: List[ShmChannel] = [] + # graph_id -> channel_id (submitter) + self._graph_owner: Dict[int, int] = {} + self._owner_lock = threading.Lock() + self._stop = threading.Event() + self._poll_thread: Optional[threading.Thread] = None + self._result_thread: Optional[threading.Thread] = None + + def setup_channels(self) -> None: + """Create shm control block + per-channel files. Idempotent w.r.t. CE + attaches — CE-side handles only need these files to exist.""" + self._ctrl = ShmControlBlock(self._server_id, create=True) + self._channels = [ + ShmChannel(self._server_id, ch_id, create=True) + for ch_id in range(self._num_channels) + ] + flexkv_logger.info( + f"TE shm dispatcher: {self._num_channels} channels created on " + f"server_id={self._server_id}" + ) + + def start_dispatch(self, transfer_manager) -> None: + """Bind the TransferManager and start the polling threads. The ctrl + ready flag is flipped so CEs that were spinning on `wait_ready` + can proceed.""" + self._tm = transfer_manager + assert self._ctrl is not None, "setup_channels() must run first" + self._ctrl.set_ready() + self._poll_thread = threading.Thread( + target=self._poll_submits, daemon=True, name="te-shm-poll" + ) + self._result_thread = threading.Thread( + target=self._poll_results, daemon=True, name="te-shm-result" + ) + self._poll_thread.start() + self._result_thread.start() + + def shutdown(self) -> None: + self._stop.set() + # Wake up futex waiters so threads can exit. + if self._ctrl is not None: + self._ctrl.notify() + for ch in self._channels: + try: + ch.close() + except Exception: + pass + self._channels = [] + if self._ctrl is not None: + try: + self._ctrl.close() + self._ctrl.unlink() + except Exception: + pass + self._ctrl = None + + def _poll_submits(self) -> None: + idle_spins = 0 + while not self._stop.is_set(): + had_work = False + for ch in self._channels: + msgs = ch.submit_recv() + if not msgs: + continue + had_work = True + for m in msgs: + if not isinstance(m, _SubmitMsg): + flexkv_logger.warning( + f"TE got unexpected submit msg type: {type(m)}" + ) + continue + graph = m.graph + with self._owner_lock: + self._graph_owner[graph.graph_id] = ch.channel_id + self._tm.submit(graph) + if had_work: + idle_spins = 0 + continue + idle_spins += 1 + if idle_spins >= 1000: + # Idle: futex wait on ctrl wake counter. + snapshot = self._ctrl.get_wake() if self._ctrl else 0 + # Re-check after snapshot — necessary to avoid lost wakeup. + any_pending = any( + ch._submit_r.value != ch._submit_w.value + for ch in self._channels + ) + if any_pending: + idle_spins = 0 + continue + if self._ctrl is not None: + self._ctrl.wait(snapshot, + timeout_ns=int(0.1 * 1_000_000_000)) + idle_spins = 0 + + def _poll_results(self) -> None: + while not self._stop.is_set(): + try: + completed = self._tm.wait(timeout=0.05) + except Exception as e: # pragma: no cover + flexkv_logger.error(f"TE result poll error: {e}") + time.sleep(0.01) + continue + if not completed: + continue + # Group completed ops by owner channel. + by_channel: Dict[int, List[CompletedOp]] = {} + for op in completed: + with self._owner_lock: + owner = self._graph_owner.get(op.graph_id) + if op.op_id == -1: + # Terminal message (completed OR failed) — drop the + # mapping after we've grouped, or failed graphs leak it. + self._graph_owner.pop(op.graph_id, None) + if owner is None: + flexkv_logger.warning( + f"TE got completed op for unknown graph {op.graph_id}" + ) + continue + by_channel.setdefault(owner, []).append(op) + for ch_id, ops in by_channel.items(): + if 0 <= ch_id < len(self._channels): + self._channels[ch_id].result_send(ops) + + +def te_shm_main(model_config: ModelConfig, + cache_config: CacheConfig, + gpu_register_port: str, + server_id: str, + num_channels: int, + start_event, + ready_event, + stop_event) -> None: + """Entrypoint for the TE subprocess in `mode="shm"`. + + Mirrors `TransferManagerInterProcessHandle._process_worker` but replaces + the single mp.Pipe with N shm channels. + + Critical ordering: shm channel files (`flexkv_te_ctrl_*`, + `flexkv_te_ch_*_*`) must exist before any CE attaches. We therefore + create the dispatcher's channels FIRST (so CEs can open the files), then + bring up the TransferManager (which blocks on GPU registration), then + bind the TM into the dispatcher and flip the ready flag. + """ + from flexkv.transfer_manager import TransferManager + dispatcher = None + tm = None + try: + os.environ["MPI4PY_RC_INITIALIZE"] = "false" + + # Phase 1: create shm channels — CE side can attach now. + dispatcher = _TEShmDispatcher(server_id, num_channels) + dispatcher.setup_channels() + # Signal start (but not ready) so the parent's `_start_event.wait()` + # returns. Ready flag is set later by start_dispatch(). + start_event.set() + + # Phase 2: build and start the TransferManager. This blocks waiting + # for GPU clients to register over the zmq gpu_register_port. + tm = TransferManager(model_config, cache_config, gpu_register_port) + tm.initialize_transfer_engine() + tm.start() + # TransferManager binds a GPU-control REP socket in __init__; every + # deployment mode must service it or a client suspend/resume call + # stalls for the full 120s RCVTIMEO. + tm.start_gpu_control_listener() + + # Phase 3: bind TM, flip ready flag, launch poll threads. + dispatcher.start_dispatch(tm) + ready_event.set() + + # Block until parent terminates us. + while not stop_event.is_set(): + stop_event.wait(timeout=1.0) + except Exception as e: + flexkv_logger.error(f"te_shm_main failed: {e}", exc_info=True) + finally: + if dispatcher is not None: + try: + dispatcher.shutdown() + except Exception: + pass + if tm is not None: + try: + tm.shutdown() + except Exception: + pass diff --git a/flexkv/transfer/transfer_engine.py b/flexkv/transfer/transfer_engine.py index fac6af941..e86b982d8 100644 --- a/flexkv/transfer/transfer_engine.py +++ b/flexkv/transfer/transfer_engine.py @@ -597,8 +597,13 @@ def _init_workers(self) -> None: for _pool in _gpu_cpu_pools: self._register_worker(_pool, TransferType.LAYERWISE, self.layerwise_workers) - if self.cache_config.enable_kv_sharing and self._cpu_handle is not None and (self.cache_config.enable_p2p_cpu \ - or (self._ssd_handle and self.cache_config.enable_p2p_ssd)): + # radixshmem mode: the CPU pool is the radix-server's SlotStore and peer + # blocks are pulled by that server (RadixClient.pull_async from the CE's + # prefetch), so FlexKV runs no peer transfer worker of its own. + if (self.cache_config.enable_kv_sharing and self._cpu_handle is not None + and not GLOBAL_CONFIG_FROM_ENV.enable_radixshmem + and (self.cache_config.enable_p2p_cpu + or (self._ssd_handle and self.cache_config.enable_p2p_ssd))): ## NOTE:if we have the cpu handle and enable p2p cpu transfer we need this worker ## (currently we inplement cpu and ssd distributed transfer in one worker) diff --git a/flexkv/transfer_manager.py b/flexkv/transfer_manager.py index 039fe9e73..9af8124a5 100644 --- a/flexkv/transfer_manager.py +++ b/flexkv/transfer_manager.py @@ -97,6 +97,9 @@ def __init__(self, self.transfer_engine: Optional[TransferEngine] = None self.storage_engine: Optional[StorageEngine] = None + # radixshmem mode: the TE's attachment to the radix-server (SlotStore = + # the CPU pool). Kept for the TE's lifetime, the pool tensors view it. + self._radix_client = None flexkv_logger.info(f"Initialized TransferManager with config successfully, " f"instance_num={self.instance_num}, expected_gpus={self.expected_gpus}") @@ -426,11 +429,20 @@ def initialize_transfer_engine(self) -> None: # KVManager; this path covers late discovery at GPU registration. recompute_cache_block_counts(self.model_config, self.cache_config) + radix_client = None + if GLOBAL_CONFIG_FROM_ENV.enable_radixshmem: + from flexkv.common.radixshmem_config import get_radixshmem_config + from flexkv.server.shm_radix_bootstrap import (attach_radix_client, + radix_index_name) + radix_client = attach_radix_client( + radix_index_name(get_radixshmem_config().local_id)) + self._radix_client = radix_client self.storage_engine = StorageEngine( self.model_config, self.cache_config, num_layers_per_pp_stage, swa_layer_groups=self.swa_layer_groups, + radix_client=radix_client, ) # Logical registration identity is separate from the CUDA device ID. @@ -577,6 +589,10 @@ def shutdown(self) -> None: # initialized manager must be safe to shut down. if getattr(self, 'transfer_engine', None) is not None: self.transfer_engine.shutdown() + if getattr(self, '_radix_client', None) is not None: + self.storage_engine = None + self._radix_client.close() + self._radix_client = None class TransferManagerOnRemote(TransferManager): """ @@ -1535,6 +1551,88 @@ def shutdown(self) -> None: flexkv_logger.info("TransferManagerMultiNodeHandle shutdown complete") +class TransferManagerShmTEProcess: + """Spawns the single TE subprocess for the multi-DP shm path. + + The bootstrap (instance 0, dp 0) creates this; CEs in other DP processes + just connect via `TransferManagerHandle(mode="shm", shm_server_id=..., + shm_channel_id=...)`. + """ + + def __init__(self, + model_config: ModelConfig, + cache_config: CacheConfig, + gpu_register_port: str, + server_id: str, + num_channels: int): + self.model_config = model_config + self.cache_config = cache_config + self.gpu_register_port = gpu_register_port + self.server_id = server_id + self.num_channels = num_channels + + self.mp_ctx = mp.get_context("spawn") + self._start_event = self.mp_ctx.Event() + self._ready_event = self.mp_ctx.Event() + self._stop_event = self.mp_ctx.Event() + self.process: Optional[Process] = None + + def start(self) -> None: + if self.process is not None and self.process.is_alive(): + return + from flexkv.transfer.shm_channel_handle import te_shm_main + # CRITICAL: clear CUDA_VISIBLE_DEVICES in the TE subprocess so it can + # cudaIpcOpenMemHandle from ALL DPs' GPUs (the parent scheduler may have + # it restricted to its own DP rank's device). mp.Process(spawn) inherits + # env from parent unless we override. Save+restore around .start(). + # + # BUT: only do this when there is genuinely more than one GPU to span + # (multi-DP/TP/CP/PP). For a single-GPU deployment (total_gpus == 1, + # e.g. vLLM serve on one restricted GPU), clearing CVD renumbers devices + # in the TE subprocess so it no longer matches the device ordinal the + # worker recorded in its TensorSharedHandle — cudaIpcOpenMemHandle then + # fails with "device >= 0 && device < num_gpus". Leave CVD untouched so + # the TE and the registering worker agree on device numbering. + clear_cvd = self.model_config.total_gpus > 1 + _saved_cuda = (os.environ.pop("CUDA_VISIBLE_DEVICES", None) + if clear_cvd else None) + try: + self.process = self.mp_ctx.Process( + target=te_shm_main, + args=(self.model_config, + self.cache_config, + self.gpu_register_port, + self.server_id, + self.num_channels, + self._start_event, + self._ready_event, + self._stop_event), + daemon=False, + ) + self.process.start() + finally: + if _saved_cuda is not None: + os.environ["CUDA_VISIBLE_DEVICES"] = _saved_cuda + self._start_event.wait() + flexkv_logger.info( + f"TransferManagerShmTEProcess started, PID={self.process.pid}, " + f"server_id={self.server_id}, channels={self.num_channels}" + ) + + def is_ready(self) -> bool: + return self._ready_event.is_set() + + def shutdown(self, timeout: float = 5.0) -> None: + if self.process is None: + return + self._stop_event.set() + self.process.join(timeout=timeout) + if self.process.is_alive(): + self.process.terminate() + self.process.join() + self.process = None + + class TransferManagerHandle: def __init__(self, model_config: ModelConfig, @@ -1563,8 +1661,22 @@ def __init__(self, self._handle: TransferManagerHandleBase = TransferManagerMultiNodeHandle( model_config, cache_config, gpu_register_port, master_host, master_ports ) + elif mode == "shm": + # Multi-DP path: each CE process gets a dedicated ShmChannel to a + # single TE subprocess. The TE is created by the bootstrap process + # via TransferManagerShmTEProcess; clients attach by server_id. + from flexkv.transfer.shm_channel_handle import ( + TransferManagerShmChannelHandle, + ) + server_id = kwargs["shm_server_id"] + channel_id = kwargs["shm_channel_id"] + self._handle: TransferManagerHandleBase = TransferManagerShmChannelHandle( + model_config, cache_config, server_id, channel_id + ) else: - raise ValueError(f"Invalid mode: {mode}, must be process, thread or remote") + raise ValueError( + f"Invalid mode: {mode}, must be process, thread, remote, or shm" + ) def start(self) -> None: self._handle.start() diff --git a/tests/test_shm_channel.py b/tests/test_shm_channel.py new file mode 100644 index 000000000..7912dea03 --- /dev/null +++ b/tests/test_shm_channel.py @@ -0,0 +1,259 @@ +"""Round-trip tests for the CE↔TE ShmChannel. + +Run with:: + + python3 -m pytest tests/test_shm_channel.py -v + +The test forks N producer processes and one consumer; each producer sends K +graph-shaped messages and reads back acks via its own result ring. +""" +from __future__ import annotations + +import multiprocessing as mp +import os +import pickle +import threading +import time + +import numpy as np +import pytest + +from flexkv.common.transfer import CompletedOp +from flexkv.transfer.shm_channel import ShmChannel, ShmControlBlock + + +SERVER_ID = "shm_channel_test" + + +def _producer(channel_id: int, n: int, server_id: str) -> None: + ch = ShmChannel(server_id, channel_id, create=False) + for i in range(n): + ch.submit_send({"channel": channel_id, "seq": i, "data": b"x" * 1024}) + # Wait for echoed acks (CompletedOps carrying channel/seq in graph_id/op_id). + received = 0 + while received < n: + msgs = ch.result_recv(timeout_s=2.0) + received += len(msgs) + for m in msgs: + assert m.graph_id == channel_id + + +def _consumer(num_channels: int, total_per_channel: int, server_id: str) -> None: + ctrl = ShmControlBlock(server_id, create=True) + channels = [ + ShmChannel(server_id, i, create=True) for i in range(num_channels) + ] + ctrl.set_ready() + pending = num_channels * total_per_channel + while pending > 0: + had_work = False + for ch in channels: + msgs = ch.submit_recv() + if msgs: + had_work = True + for m in msgs: + # Echo back the (channel, seq) as a CompletedOp — the result + # ring is now typed to CompletedOp records. + ch.result_send([CompletedOp(graph_id=m["channel"], + op_id=m["seq"])]) + pending -= 1 + if not had_work: + time.sleep(0.001) + for ch in channels: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def _cleanup_shm(server_id: str, num_channels: int) -> None: + for ch_id in range(num_channels): + path = f"/dev/shm/flexkv_te_ch_{server_id}_{ch_id}" + if os.path.exists(path): + os.unlink(path) + ctrl_path = f"/dev/shm/flexkv_te_ctrl_{server_id}" + if os.path.exists(ctrl_path): + os.unlink(ctrl_path) + + +def test_n_producer_one_consumer(): + server_id = f"{SERVER_ID}_npm" + num_channels = 4 + total_per_channel = 32 + + _cleanup_shm(server_id, num_channels) + + ctx = mp.get_context("spawn") + consumer = ctx.Process( + target=_consumer, + args=(num_channels, total_per_channel, server_id), + ) + consumer.start() + + # Wait for consumer to set up shm files. + deadline = time.monotonic() + 5.0 + while time.monotonic() < deadline: + if os.path.exists(f"/dev/shm/flexkv_te_ctrl_{server_id}"): + break + time.sleep(0.01) + else: + consumer.terminate() + pytest.fail("consumer never created control block") + + # Wait for ready flag via control block. + ctrl = ShmControlBlock(server_id, create=False) + assert ctrl.wait_ready(timeout_s=5.0) + ctrl.close() + + producers = [ + ctx.Process( + target=_producer, args=(i, total_per_channel, server_id) + ) + for i in range(num_channels) + ] + for p in producers: + p.start() + for p in producers: + p.join(timeout=10.0) + assert p.exitcode == 0, f"producer {p.pid} exit code {p.exitcode}" + + consumer.join(timeout=10.0) + assert consumer.exitcode == 0 + + +def test_single_round_trip_local(): + server_id = f"{SERVER_ID}_local" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + ch = ShmChannel(server_id, 0, create=True) + try: + ch.submit_send("hello") + ch.submit_send({"k": 42}) + msgs = ch.submit_recv() + assert msgs == ["hello", {"k": 42}] + + # Result ring carries fixed-width CompletedOp records; all fields must + # round-trip, including the transfer_type string and the -1 sentinel. + sent = [ + CompletedOp(graph_id=7, op_id=3, transfer_type="H2D", + num_blocks=12, num_bytes=98304), + CompletedOp(graph_id=7, op_id=-1), # graph-completed sentinel + ] + ch.result_send(sent) + out = ch.result_recv(timeout_s=0.0) + assert out == sent + assert out[1].is_graph_completed() + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def test_submit_fragmentation(): + """A payload larger than one slot must fragment and round-trip intact, + interleaved with small single-slot messages.""" + server_id = f"{SERVER_ID}_frag" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + # Small slots so a modest payload spans many fragments. + ch = ShmChannel(server_id, 0, create=True, + submit_slots=1024, slot_size=4096) + try: + big = {"arr": np.arange(200_000, dtype=np.int64)} # ~1.5 MB > slot + small = {"k": 1} + ch.submit_send(small) + ch.submit_send(big) + ch.submit_send(small) + msgs = ch.submit_recv() + assert len(msgs) == 3 + assert msgs[0] == small + assert np.array_equal(msgs[1]["arr"], big["arr"]) + assert msgs[2] == small + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def test_submit_payload_too_large(): + """A payload that can't fit the whole ring is rejected, not deadlocked.""" + server_id = f"{SERVER_ID}_toobig" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + ch = ShmChannel(server_id, 0, create=True, + submit_slots=8, slot_size=4096) + try: + with pytest.raises(ValueError): + ch.submit_send(b"x" * (8 * 4096)) # needs more fragments than slots + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def _make_big_graph(nbytes: int) -> dict: + """A graph-shaped payload that pickles to ~nbytes (block-id arrays dominate).""" + n = nbytes // 16 # two int64 arrays + return {"src": np.arange(n, dtype=np.int64), + "dst": np.arange(n, dtype=np.int64)} + + +def test_submit_big_graph_500kb(): + """A ~500 KB graph fragments across the default 32 KB slots and round-trips.""" + server_id = f"{SERVER_ID}_big500" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + ch = ShmChannel(server_id, 0, create=True) # default 32 KB / 8192 + try: + big = _make_big_graph(500 * 1024) + blob_sz = len(pickle.dumps(big, protocol=pickle.HIGHEST_PROTOCOL)) + assert blob_sz > 500 * 1024, f"payload only {blob_sz} B" + assert blob_sz > ch.slot_size, "payload must exceed one slot" + ch.submit_send(big) + out = ch.submit_recv() + assert len(out) == 1 + assert np.array_equal(out[0]["src"], big["src"]) + assert np.array_equal(out[0]["dst"], big["dst"]) + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def test_submit_small_graphs_high_rate(): + """4096 small multi-slot graphs at ~4096/s: a producer thread submits while a + consumer thread drains, verifying no loss, correct order, and no ring-full + stall at the default 8192-slot capacity.""" + server_id = f"{SERVER_ID}_hirate" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + ch = ShmChannel(server_id, 0, create=True) # default 32 KB / 8192 + n_msgs = 4096 + small = {"payload": b"x" * (2 * ch.slot_size)} # spans ~3 slots each + received: list = [] + + def consumer() -> None: + while len(received) < n_msgs: + received.extend(ch.submit_recv()) + + try: + t = threading.Thread(target=consumer, daemon=True) + t.start() + start = time.monotonic() + for i in range(n_msgs): + ch.submit_send({"seq": i, **small}) + t.join(timeout=30.0) + elapsed = time.monotonic() - start + assert len(received) == n_msgs, f"got {len(received)}/{n_msgs}" + assert [m["seq"] for m in received] == list(range(n_msgs)), "order/loss" + rate = n_msgs / elapsed + assert rate >= 4096, f"throughput {rate:.0f}/s below 4096/s target" + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() From 87c02b8f6f8cb3b51e1efa98684f0d27dc80f4c8 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:26:42 +0800 Subject: [PATCH 02/21] hash: feed numpy buffers to the hasher directly instead of torch.from_numpy torch.from_numpy is not safe to call concurrently from several planner threads. c_ext gains Hasher.update_numpy and gen_hashes_numpy, which read the pybind11 buffer directly, and flexkv.common.hash_utils uses them. Part 2/7 of the radixshmem rebase (see part 1 for provenance). --- csrc/bindings.cpp | 30 ++++++++++++++++++++++++++++++ flexkv/common/hash_utils.py | 6 +++--- 2 files changed, 33 insertions(+), 3 deletions(-) diff --git a/csrc/bindings.cpp b/csrc/bindings.cpp index d7a33b19a..64f94101f 100644 --- a/csrc/bindings.cpp +++ b/csrc/bindings.cpp @@ -10,6 +10,7 @@ #include #include #include +#include #include #include #include @@ -481,6 +482,24 @@ PYBIND11_MODULE(c_ext, m) { m.def("gen_hashes", &flexkv::gen_hashes, "Generate hashes for a tensor", py::arg("hasher"), py::arg("token_ids"), py::arg("tokens_per_block"), py::arg("block_hashes")); + m.def( + "gen_hashes_numpy", + [](flexkv::Hasher &hasher, py::array token_ids, int tokens_per_block, + py::array block_hashes) { + // numpy-buffer variant of gen_hashes; bypasses torch.from_numpy. + py::buffer_info tok = token_ids.request(); + py::buffer_info bh = block_hashes.request(true); + const int64_t *tok_ptr = static_cast(tok.ptr); + flexkv::HashType *bh_ptr = static_cast(bh.ptr); + for (py::ssize_t i = 0; i < bh.size; i++) { + hasher.update(tok_ptr + i * tokens_per_block, + tokens_per_block * sizeof(int64_t)); + bh_ptr[i] = hasher.digest(); + } + }, + "Generate block hashes directly from numpy buffers", py::arg("hasher"), + py::arg("token_ids"), py::arg("tokens_per_block"), + py::arg("block_hashes")); py::class_(m, "SSDIOCTX") .def( @@ -766,6 +785,17 @@ PYBIND11_MODULE(c_ext, m) { py::overload_cast(&flexkv::Hasher::update), "Update the hasher with pointer and size", py::arg("input"), py::arg("size")) + .def( + "update_numpy", + [](flexkv::Hasher &self, py::array arr) { + // Hash the numpy buffer directly, bypassing torch.from_numpy + // (whose tensor conversion is not concurrency-safe). + py::buffer_info info = arr.request(); + self.update(info.ptr, + static_cast(info.size * info.itemsize)); + }, + "Update the hasher directly from a numpy array buffer", + py::arg("input")) .def("digest", &flexkv::Hasher::digest, "Return the hash value"); #ifdef FLEXKV_ENABLE_CFS py::class_(m, "Pcfs") diff --git a/flexkv/common/hash_utils.py b/flexkv/common/hash_utils.py index a625bb1f6..1f0df25bc 100644 --- a/flexkv/common/hash_utils.py +++ b/flexkv/common/hash_utils.py @@ -2,7 +2,6 @@ from typing import NewType, Optional import numpy as np -import torch from flexkv import c_ext @@ -20,7 +19,8 @@ def reset(self) -> None: self.hasher.reset() def update(self, array: np.ndarray) -> None: - self.hasher.update(torch.from_numpy(array)) + # update_numpy avoids torch.from_numpy, which is not concurrency-safe. + self.hasher.update_numpy(np.ascontiguousarray(array)) def digest(self) -> HashType: return HashType(self.hasher.digest()) @@ -44,7 +44,7 @@ def gen_hashes(token_ids: np.ndarray, tokens_per_block: int, hasher: Optional[Ha block_hashes = np.zeros(token_ids.size // tokens_per_block, dtype=np.uint64) if hasher is None: hasher = Hasher() - c_ext.gen_hashes(hasher.hasher, torch.from_numpy(token_ids), tokens_per_block, torch.from_numpy(block_hashes)) + c_ext.gen_hashes_numpy(hasher.hasher, np.ascontiguousarray(token_ids), tokens_per_block, block_hashes) return block_hashes if __name__ == "__main__": From e7c3f59b768d02c11039723d4550c84275178033 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:26:42 +0800 Subject: [PATCH 03/21] radixshmem: one global YAML config and the radix-server bootstrap * flexkv/common/radixshmem_config.py: RadixShmemConfig, loaded once from FLEXKV_RADIXSHMEM_CONFIG_PATH; server/client/index/distributed sections, index.register_chunk_size defaulting to 4096 / tokens_per_block. * flexkv/common/config.py, flexkv/integration/config.py: FLEXKV_ENABLE_RADIXSHMEM / enable_radixshmem, node-local DP (ModelConfig.local_dp_size, RankInfo.local_dp_client_id) for every tier, and the sglang dp_rank check under radixshmem with dp_size > 1. * flexkv/server/shm_radix_bootstrap.py: create the radix regions and the embedded radix-server subprocess; size the CPU FULL / SWA slot pools from the cache config so the SlotStore stride equals FlexKV's block. * examples/radixshmem_configs/*.yaml, docs/radixshmem/config_zh.md. Part 3/7 of the radixshmem rebase (see part 1 for provenance). Co-authored-by: Hao Xu Co-authored-by: Iris Ge Co-authored-by: linhu-nv Co-authored-by: teeebin --- docs/radixshmem/config_zh.md | 272 ++++++++++ .../radixshmem_multi_node.yaml | 24 + .../radixshmem_single_node.yaml | 17 + flexkv/common/config.py | 63 +++ flexkv/common/radixshmem_config.py | 365 +++++++++++++ flexkv/integration/config.py | 113 +++- flexkv/server/shm_radix_bootstrap.py | 483 ++++++++++++++++++ 7 files changed, 1334 insertions(+), 3 deletions(-) create mode 100644 docs/radixshmem/config_zh.md create mode 100644 examples/radixshmem_configs/radixshmem_multi_node.yaml create mode 100644 examples/radixshmem_configs/radixshmem_single_node.yaml create mode 100644 flexkv/common/radixshmem_config.py create mode 100644 flexkv/server/shm_radix_bootstrap.py diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md new file mode 100644 index 000000000..309117c9b --- /dev/null +++ b/docs/radixshmem/config_zh.md @@ -0,0 +1,272 @@ +# radixshmem 模式配置参考 + +本文列出 FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)时的全部配置项: +哪些走环境变量、哪些走 YAML、哪些由 FlexKV 自己推导而禁止手工设置。 + +配置分三层: + +| 层 | 载体 | 内容 | +|---|---|---| +| 进程级开关 | 环境变量 `FLEXKV_RADIX_*` | 是否启用、YAML 路径、server 启动方式,以及两个仅供同机多节点测试的 per-node 覆盖 | +| 集群配置 | YAML,`FLEXKV_RADIXSHMEM_CONFIG_PATH` 指向 | 全局,所有节点逐字节相同;键名与 radixshmem 的 dataclass 字段一致 | +| 几何 | 由 `ModelConfig` / `CacheConfig` 推导 | slot 数、slot 字节数、对齐、shm 名;不可配置 | + +同一个值只有一个来源。YAML 里没写的键取本文列出的默认值;没有 YAML 时全部取默认值,即单机模式。 +示例文件在 `examples/radixshmem_configs/`:`radixshmem_single_node.yaml` 和 `radixshmem_multi_node.yaml`。 +实现在 `flexkv/common/radixshmem_config.py`。 + +--- + +## 1. 环境变量 + +| 变量 | 默认 | 说明 | +|---|---|---| +| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担,KVServer 不启动,每个 DP 进程各建一个 KVTaskEngine 并 attach 共享的 radix 区域。在 `flexkv` 首次 import 前设置。 | +| `FLEXKV_RADIXSHMEM_CONFIG_PATH` | 空 | 第 2 节 YAML 的路径。为空时所有键取默认值。 | +| `FLEXKV_RADIX_SERVER_LAUNCH_MODE` | `embedded` | `embedded`:dp0 进程以子进程方式启动 radix-server;`external`:attach 运维已启动的 radix-server。 | +| `FLEXKV_RADIX_NODE_NAME` | 空 | per-node 覆盖,见 3.3。生产部署不设。 | +| `FLEXKV_RADIX_RPC_ADDRESS` | 空 | per-node 覆盖,见 3.3。设了就忽略 YAML 的 `cluster.rpc_interface`。生产部署不设。 | + +另有两个 FlexKV 通用变量在该模式下有约束: + +- `FLEXKV_CPU_LAYOUT` 必须是 `BLOCKFIRST`。一个 SlotStore slot 就是一个连续的 CPU block,LAYERFIRST 给不出这个布局。 +- `FLEXKV_HUGETLBFS_DIR`(默认 `/mnt/hugepages`):`server.hugepage_path` 为空且 `CacheConfig.use_hugepage_cpu_buffer` 为真时,radix 区域建在这个 hugetlbfs 挂载点下。 + +该模式与 `enable_ssd`、`enable_remote` 互斥,启动时报错。`enable_p2p_cpu` / `enable_p2p_ssd` 也必须为 False: +跨节点复用由 radix-server 自己完成(etcd + RDMA),在 YAML 使集群成为分布式(`expected_min_nodes > 1` 或 +`num_rht_shards > 1`)时自动开启,不再经过 FlexKV 的 Redis P2P 路径。 + +--- + +## 2. YAML 字段 + +五个段。`cluster` / `data` / `index` / `server` 四段的键按名字直接构造 radixshmem 的 +`ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig`,用 +`dataclasses.fields()` 校验:未知键报错,几何键(2.6)报错。`client` 段是 FlexKV 自己的参数。 + +### 2.1 `cluster`(`shmradix.ClusterConfig`) + +| 键 | 默认 | 说明 | +|---|---|---| +| `cluster_id` | `flexkv` | 集群命名空间。既是 etcd 键前缀 `radix//...`,也派生本机全部 shm 和 socket 名(见 4)。同一台机器上跑多个 FlexKV 实例时用不同的 `cluster_id` 区分。 | +| `expected_min_nodes` | `0` | 集群开关。`> 1` 时进入 etcd + RDMA 模式:bootstrap 等到 etcd 里登记的节点数达到 `max(expected_min_nodes, num_rht_shards)` 且稳定 `settle_ms` 后分配 rank。`world_size` 是实际观察到的节点数,可以大于该值。 | +| `registry` | `etcd://127.0.0.1:2379` | etcd 地址,集群模式必填。格式 `etcd://host:port`,多个成员在 scheme 之后用逗号或分号分隔,scheme 只写一次:`etcd://10.0.0.1:2379,10.0.0.2:2379`。见 3.4。 | +| `rpc_interface` | 空 | 解析 bootstrap IP 的网卡名,如 `bond0`(南北向管理网卡)。集群模式下必填(除非用 `FLEXKV_RADIX_RPC_ADDRESS` 覆盖)。每个节点解析出自己的 IP,节点身份自动派生为 `node`。 | +| `rpc_port` | `0` | bootstrap / XRC 监听端口,0 由系统分配。 | +| `settle_ms` | `500` | 成员集合稳定多久后开始 bootstrap。 | +| `bootstrap_timeout_sec` | `120` | rendezvous 超时。FlexKV 的 attach 超时是该值加 60 秒。 | +| `index_dev` | 空 | index 控制面(RHT 面和 peer index 面)用的 HCA,空为第一个可用设备,通常是 `mlx5_0` 即东西向计算网卡。建议指定南北向管理网卡的 HCA(如 `mlx5_bond_0`):控制面只有小消息,计算网卡留给 KV 字节。KV 字节的 HCA 是 `data.transfer_devices`。 | +| `gid_idx` | `3` | RoCE GID 索引。 | +| `rht_transport` | `xrc` | client 到 RHT shard holder 的 QP 类型,`xrc` 或 `dc`。 | +| `peer_index_transport` | `xrc` | remote walk 读 peer index 的 QP 类型,`xrc` 或 `dc`。 | +| `remote_op_transport` | `zmq` | remote insert / query 控制面,`zmq` 或 `dc`。FlexKV 不开 remote op,该字段不生效。 | +| `num_rht_shards` | `0` | RHT 分片数。0 为每节点一片。设了必须不大于节点数。 | +| `rht_shard_holders` | `[]` | 持有分片的 rank 列表,空为 rank 0 到 `num_rht_shards - 1`。rank 由 rendezvous 后按 `node` 字典序分配。 | +| `rht_slots_per_bucket` | `4` | RHT 每 bucket 的 slot 数,取 1 / 2 / 4 / 8。1 是盲覆盖,会丢路由项。 | +| `enable_remote_insert` | `false` | 透传。 | +| `enable_remote_query` | `false` | 透传。 | +| `zmq_listen_port` | `0` | 透传。 | + +`rht_transport` / `peer_index_transport` / `remote_op_transport` / `num_rht_shards` / `rht_shard_holders` / +`rht_slots_per_bucket` 只有 rank 0 的值生效,bootstrap 时经 etcd `/config` 广播给其他节点。全局 YAML +下各节点值本来相同,这一规则只是多一层保险。 + +禁止出现:`node_name`、`rpc_address`。它们是 per-node 值,全局 YAML 放不下;需要时走 3.3 的环境变量。 + +### 2.2 `data`(`shmradix.DataPlaneConfig` 的非几何字段) + +| 键 | 默认 | 说明 | +|---|---|---| +| `transfer_devices` | `[]` | KV 字节传输(mooncake)用的 HCA 列表,空为 mooncake 发现的全部设备。 | +| `transfer_protocol` | `rdma` | `rdma` 或 `tcp`。 | +| `transfer_ip` | 空 | 数据面 IP,空为 rpc 地址。 | +| `transfer_port` | `0` | 0 由引擎选。 | +| `transfer_metadata` | `P2PHANDSHAKE` | mooncake 元数据服务。 | +| `prefault` | `true` | server 启动时 MAP_POPULATE 整个 SlotStore。D2H 延迟可预测,页由 server 进程的 NUMA 策略放置;大池会拉长启动时间。 | +| `max_inflight` | `256` | 传输引擎在飞 batch 数。 | +| `max_pending_jobs` | `4096` | 排队 + 运行 + 未领取的 job 上限,超过则 Submit 被拒。 | +| `job_ttl_s` | `60.0` | 未领取 job 的保留秒数。 | + +禁止出现:`data_bytes`、`full_slot_bytes`、`swa_slot_bytes`、`mamba_slot_bytes`、`slot_align`、`data_name`。 + +### 2.3 `index`(`shmradix.IndexConfig` 的非几何字段) + +| 键 | 默认 | 说明 | +|---|---|---| +| `data_pool_ratio` | `8.0` | 索引 DataPool 大小系数:`full_slots × ratio × (12 或 16)` 字节。 | +| `background_evict_ratio` | `0.05` | 后台驱逐比例,0 关闭。 | +| `max_nodes` | `0` | radix 节点池容量,0 自动。 | +| `register_chunk_size` | `4096 / tokens_per_block` | RHT 注册粒度(block 数)。FlexKV 的默认让一段覆盖 4096 个 token,与 block 大小无关(radixshmem 自身默认 128 block)。 | + +禁止出现:`name`、`tokens_per_block`、`full_slots`、`swa_slots`、`swa_window_blocks`、`mamba_slots`、`evict_policy`。 + +### 2.4 `server`(`shmradix.RadixServerConfig` 顶层) + +| 键 | 默认 | 说明 | +|---|---|---| +| `endpoint` | 空 | gRPC 端点。空为 `unix:///dev/shm/.sock`。server 监听和 client attach 都用它。 | +| `rpc_workers` | `32` | gRPC 工作线程数。每个有在飞 job 的 client 占一个。 | +| `hugepage_path` | 空 | index 和 SlotStore 的 hugetlbfs 挂载点。空时按第 1 节 `FLEXKV_HUGETLBFS_DIR` 的规则决定。 | + +### 2.5 `client`(FlexKV 侧,不传给 radixshmem) + +| 键 | 默认 | 说明 | +|---|---|---| +| `prefetch_timeout_ms` | `5000` | 一次 prefetch 拉取的服务端超时。到期后 job 以本地命中的部分完成。 | +| `prefetch_max_inflight` | `128` | 每个 DP 进程在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | +| `max_outstanding` | `256` | `RadixClient` 未领取 job 的上限。 | + +### 2.6 由 FlexKV 推导、禁止手工设置的字段 + +| 字段 | 来源 | +|---|---| +| `index.name` | `/shmradix__cpu` | +| `index.tokens_per_block` | `CacheConfig.tokens_per_block` | +| `index.full_slots` | `CacheConfig.num_cpu_blocks` | +| `index.swa_slots` / `swa_window_blocks` | `CacheConfig.swa.num_slots` / `window_blocks`,SWA 未开启为 0 | +| `data.full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 block 字节数 | +| `data.swa_slot_bytes` | 同上,SWA 池按 uint8 | +| `data.data_bytes` | `full_slots × full_slot_bytes + swa_slots × swa_slot_bytes` | +| `data.slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,保证 slot stride 等于 block 大小 | +| `data.data_name` | `_data` | + +这些值在 YAML 中出现时启动报错,防止和 `CacheConfig` 静默冲突。TE 进程 attach 后还会用 +`check_geometry` 复核 server 端区域和 FlexKV 自己的布局一致。 + +--- + +## 3. 示例 + +### 3.1 单机 + +不设 `FLEXKV_RADIXSHMEM_CONFIG_PATH` 即可。等价于: + +```yaml +cluster: + cluster_id: flexkv + expected_min_nodes: 0 +``` + +无 etcd、无 RDMA 依赖。多个 DP 进程共享一个 radix-server 和一个 TE。 + +### 3.2 多机(全局配置,所有节点同一文件) + +```yaml +# /etc/flexkv/radixshmem.yaml +cluster: + cluster_id: prod_a + expected_min_nodes: 4 + num_rht_shards: 4 + registry: etcd://10.0.0.1:2379 + rpc_interface: bond0 # 南北向网卡;每节点解析自己的 IP,身份为 node + index_dev: mlx5_bond_0 # index 内部 RDMA 的 HCA,南北向网卡 + gid_idx: 3 + rht_transport: xrc + peer_index_transport: dc + rht_slots_per_bucket: 4 + bootstrap_timeout_sec: 120 +data: + transfer_devices: [mlx5_1, mlx5_2] # KV 字节传输的 HCA + prefault: true +index: + data_pool_ratio: 8.0 +server: + rpc_workers: 32 +client: + prefetch_timeout_ms: 5000 + prefetch_max_inflight: 128 +``` + +每个节点: + +```bash +export FLEXKV_ENABLE_RADIXSHMEM=1 +export FLEXKV_RADIXSHMEM_CONFIG_PATH=/etc/flexkv/radixshmem.yaml +export FLEXKV_CPU_LAYOUT=BLOCKFIRST +``` + +节点身份、shm 名、rank 全部自动派生,文件里没有任何 per-node 内容。 + +### 3.3 同机多节点(测试) + +两个 radix-server 在一台机器上时,同一网卡解析出同一 IP,身份会撞。用两个 per-node 环境变量区分: + +```bash +# 进程 A +FLEXKV_RADIX_NODE_NAME=r0 FLEXKV_RADIX_RPC_ADDRESS=127.0.0.1 ... +# 进程 B +FLEXKV_RADIX_NODE_NAME=r1 FLEXKV_RADIX_RPC_ADDRESS=127.0.0.1 ... +``` + +`FLEXKV_RADIX_RPC_ADDRESS` 设置后 FlexKV 清掉 YAML 的 `rpc_interface`(radixshmem 规则是 interface 优先, +不清会被覆盖回去)。两个进程仍共用同一份 YAML。 + +设置了 `FLEXKV_RADIX_NODE_NAME` 时,本机命名前缀从 `` 变为 `_`(第 4 节), +两个进程的 SlotStore、socket 和 TE channel 因此互不冲突。每个进程还要各给一个 `FLEXKV_SERVER_RECV_PORT`。 + +### 3.4 `registry` 的填法 + +radixshmem 把 `registry` 去掉第一个 `://` 之前的 scheme 后,余下部分按逗号或分号切成 endpoint 列表, +交给 etcd 的 clientv3。因此: + +- 单成员:`etcd://10.0.0.1:2379`。 +- 多成员:`etcd://10.0.0.1:2379,10.0.0.2:2379,10.0.0.3:2379`。scheme 只写一次;写成 + `etcd://a:2379,etcd://b:2379` 会把第二个 `etcd://b:2379` 原样当作 endpoint 传下去,连接失败。 +- 只支持明文连接,没有 TLS 和用户名密码的配置入口。拨号超时固定 5 秒。 +- 一个 etcd 可以服务多个集群,键空间由 `cluster_id` 隔开(`radix//...`);索引 rendezvous 和 + 数据面登记(`data/`)都在同一个 etcd 里。 +- etcd 不只在启动时用:节点的 lease keep-alive、`/peers` watch 和数据面登记贯穿整个运行期,etcd 不可用会导致 + lease 过期、节点从集群视图中消失。生产环境用 3 成员 etcd,并把全部成员写进 `registry`。 +- 每个节点必须能访问 `registry` 里的地址;默认值 `127.0.0.1:2379` 只适用于所有节点在同一台机器上的测试。 +- 同一进程内 etcd 连接是全局单例,首个 `init` 的 endpoint 生效;FlexKV 里索引和数据面都用 `cluster.registry`, + 不会出现两个不同地址。 + +--- + +## 4. 命名派生 + +所有名字来自 `cluster.cluster_id`。记本机前缀 `local_id`:未设 `FLEXKV_RADIX_NODE_NAME` 时就是 +`cluster_id`,设了则是 `_`(`RadixShmemConfig.local_id`)。 + +| 对象 | 名字 | +|---|---| +| etcd 键空间 | `radix//...` | +| index shm | `/shmradix__cpu`;集群模式下 radixshmem 再追加 `_`,attach 方只需 base name | +| SlotStore shm | `/shmradix__cpu_data` | +| gRPC socket | `/dev/shm/shmradix__cpu.sock` | +| TE shm channel | FlexKV 内部 IPC 名,以 `local_id` 为前缀 | + +--- + +## 5. 三套传输的区分 + +| 配置 | 取值 | 链路 | HCA | +|---|---|---|---| +| `cluster.rht_transport` | xrc / dc | client 向 RHT shard holder 写路由项 | `cluster.index_dev` | +| `cluster.peer_index_transport` | xrc / dc | remote walk 时对 peer 节点 index 的单边 RDMA read | `cluster.index_dev` | +| `cluster.remote_op_transport` | zmq / dc | remote insert / query 控制面,FlexKV 不使用 | zmq 走 TCP | +| `data.transfer_protocol` + `data.transfer_devices` | rdma / tcp | 两节点 SlotStore 之间的 KV 字节搬运(mooncake),即 `pull_async` 的实际拉取 | `data.transfer_devices` | + +xrc 对每个目标一条 QP;dc 用一个 DC initiator 对所有目标,QP 数 O(1),需要 mlx5。任一 index 面为 dc 时 +server 建 DCT,client 两个面共用一个 DCI。FlexKV 的 `get_match` 只查本地,不走前两条;`pull_async` +在服务端规划时走 RHT 和 remote walk,随后的字节搬运走 mooncake。 + +--- + +## 6. 启动时校验 + +FlexKV 在加载 YAML 时检查以下条件,不满足直接报错,不等到 rendezvous: + +- 未知键、2.6 的几何键、`cluster.node_name`、`cluster.rpc_address` 出现在 YAML。 +- `expected_min_nodes > 1` 时 `registry` 为空,或 `rpc_interface` 与 `FLEXKV_RADIX_RPC_ADDRESS` 都为空。 +- `FLEXKV_RADIX_RPC_ADDRESS=0.0.0.0`:所有节点会派生出同一个身份。 +- `rht_shard_holders` 和 `transfer_devices` 接受列表或逗号分隔字符串,其他类型报错。 +- `num_rht_shards > expected_min_nodes`(`expected_min_nodes` 非 0 时)。 +- `rht_slots_per_bucket` 不在 {1, 2, 4, 8}。 +- `rht_transport` / `peer_index_transport` 不在 {xrc, dc};`remote_op_transport` 不在 {zmq, dc}。 +- `client.prefetch_max_inflight >= client.max_outstanding`。 +- `FLEXKV_CPU_LAYOUT != BLOCKFIRST`,或 `num_cpu_blocks <= 0`,或 SWA 池装不下一个 window。 +- `CacheConfig` 打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`。 + +TE attach 后 `check_geometry` 复核 server 端的 `tokens_per_block`、各池 slot 数、slot 字节数和 stride, +不一致则报错退出,不会静默错位传输。 diff --git a/examples/radixshmem_configs/radixshmem_multi_node.yaml b/examples/radixshmem_configs/radixshmem_multi_node.yaml new file mode 100644 index 000000000..ca19e731a --- /dev/null +++ b/examples/radixshmem_configs/radixshmem_multi_node.yaml @@ -0,0 +1,24 @@ +# FlexKV radixshmem mode, a 4-node cluster sharing CPU KV over RDMA. +# +# export FLEXKV_ENABLE_RADIXSHMEM=1 +# export FLEXKV_RADIXSHMEM_CONFIG_PATH=/etc/flexkv/radixshmem.yaml +# export FLEXKV_CPU_LAYOUT=BLOCKFIRST +# +# This file is GLOBAL: copy it byte for byte to every node. Each node resolves +# its own IP from cluster.rpc_interface and gets its identity and rank from the +# etcd rendezvous. Only the keys a cluster must set are listed; everything else +# keeps its default. Reference: docs/radixshmem/config_zh.md + +cluster: + cluster_id: prod_a + expected_min_nodes: 4 + num_rht_shards: 4 + registry: etcd://10.0.0.1:2379 + rpc_interface: bond0 + index_dev: mlx5_bond_0 + +data: + transfer_devices: [mlx5_0, mlx5_1, mlx5_2, mlx5_3, mlx5_4, mlx5_5, mlx5_6, mlx5_7] + +index: + background_evict_ratio: 0.05 diff --git a/examples/radixshmem_configs/radixshmem_single_node.yaml b/examples/radixshmem_configs/radixshmem_single_node.yaml new file mode 100644 index 000000000..2bc9ba649 --- /dev/null +++ b/examples/radixshmem_configs/radixshmem_single_node.yaml @@ -0,0 +1,17 @@ +# FlexKV radixshmem mode, one node (no etcd, no RDMA). +# +# export FLEXKV_ENABLE_RADIXSHMEM=1 +# export FLEXKV_RADIXSHMEM_CONFIG_PATH=$PWD/examples/radixshmem_configs/radixshmem_single_node.yaml +# export FLEXKV_CPU_LAYOUT=BLOCKFIRST +# +# Every key below is at its default; the file is equivalent to setting no +# FLEXKV_RADIXSHMEM_CONFIG_PATH at all. It is here to show what is tunable on a +# single node. Slot counts / slot bytes / shm names are derived from the FlexKV +# cache configuration and must not appear here. Reference: docs/radixshmem/config_zh.md + +cluster: + cluster_id: flexkv + expected_min_nodes: 0 + +index: + background_evict_ratio: 0.05 diff --git a/flexkv/common/config.py b/flexkv/common/config.py index 5573fef2a..2ddc307c7 100644 --- a/flexkv/common/config.py +++ b/flexkv/common/config.py @@ -185,6 +185,12 @@ class ModelConfig: # and token_size_in_bytes/num_cpu_blocks are computed by summing across groups. layer_groups: Optional[List[LayerGroupSpec]] = None + # SGLang DP-Attention node-local width. Set when every DP group lives on + # one node, which makes FlexKV form one instance per node instead of one + # spanning the cluster (FlexKVConfig.get_sglang_node_local_dp_size). + # None keeps the cross-node path. + local_dp_size: Optional[int] = None + # ------------------------------------------------------------------ # Freeze mechanism: after post_init, ModelConfig must not be mutated # ------------------------------------------------------------------ @@ -207,6 +213,24 @@ def freeze(self) -> None: f"[ModelConfig] cannot derive gpus_per_node: " f"total_gpus={self.total_gpus} not divisible by nnodes={self.nnodes}" ) + if self.local_dp_size is not None: + if self.local_dp_size < 1: + raise ValueError( + "[ModelConfig] local_dp_size must be >= 1, got " + f"{self.local_dp_size}" + ) + if not self.enable_dp_attention or self.pp_size != 1: + raise ValueError( + "[ModelConfig] local_dp_size is only supported for " + "SGLang DP Attention with pp_size=1" + ) + if self.nnodes * self.local_dp_size != self.dp_size: + raise ValueError( + "[ModelConfig] node-local DP requires every DP group to " + "reside on exactly one node, but " + f"nnodes={self.nnodes} * local_dp_size={self.local_dp_size} " + f"!= dp_size={self.dp_size}" + ) if self.nnodes_per_pp_rank > 2: raise ValueError( f"[ModelConfig] only support 2-nodes TP for now, but got " @@ -512,6 +536,15 @@ def dp_client_id(self) -> int: """ return self.instance_id * self.model_config.dp_size + self.dp_rank + @property + def local_dp_client_id(self) -> int: + """Dense node-local id: which rank owns this node's KVServer or + radix-server and its TE channel, and the shared-memory IPC names.""" + local_dp_size = self.model_config.local_dp_size + if local_dp_size is None: + return self.dp_client_id + return self.dp_rank % local_dp_size + @property def attn_tp_rank(self) -> int: """Compatibility alias for the normalized attention TP rank.""" @@ -589,6 +622,10 @@ def __str__(self) -> str: f", local_rank={self.local_rank}, effective_tp_rank={self.effective_tp_rank}" ) + +RADIX_SWA_WINDOW_BLOCKS = 8 + + @dataclass class SWAPoolConfig: """Configuration for SWA (Sliding Window Attention) host pool(s). @@ -603,11 +640,17 @@ class SWAPoolConfig: num_remote_slots: int = 0 # Number of REMOTE SWA pool slots (0 = no REMOTE SWA tier) num_swa_layers: int = 61 # Number of SWA layers (all 61 for DSv4) bytes_per_token_per_layer: int = 584 # nope_fp8(448) + rope_bf16(128) + scale(8) + + window_blocks: int = RADIX_SWA_WINDOW_BLOCKS # True when the SWA page also carries heterogeneous sidecar groups (for # example DeepSeek-V4 attention/indexer compress states). multi_group: bool = False evict_ratio: float = 0.1 # Fraction of pool to evict when full pin_memory: bool = True # Use pinned memory for async DMA + # Sidecar groups packed into one SWA page (DSv4 compress states). The TE + # learns them from the GPU registration; the radixshmem bootstrap needs them + # earlier to size the SWA slot, so the connector records them here. + layer_groups: Optional[List['LayerGroupSpec']] = None def for_ssd_tier(self) -> "SWAPoolConfig": """Derive the SSD-tier SWA config (same slot geometry, num_ssd_slots slots). @@ -724,6 +767,9 @@ class CacheConfig: # Stored for deferred recomputation when layer_groups become known _user_cpu_cache_gb: float = 0 _user_ssd_cache_gb: float = 0 + # Layers one CPU block covers (this node's PP stage), recorded by the + # adapters' config resolution; 0 = derive num_layers // pp_size. + _num_layers_per_pp_stage: int = 0 # SWA pool config (DeepSeek V4) swa: Optional['SWAPoolConfig'] = None @@ -793,6 +839,22 @@ def __str__(self) -> str: server_launch_mode=os.getenv('FLEXKV_SERVER_LAUNCH_MODE', 'embedded').lower(), server_recv_port=os.getenv('FLEXKV_SERVER_RECV_PORT', 'ipc:///tmp/flexkv_server'), + # radixshmem mode: the CPU tier is radixshmem's index + SlotStore, one + # radix-server per node, one shared TE, a KVTaskEngine per DP process (no + # KVServer). Everything else about that mode -- cluster membership, RDMA + # devices, prefetch limits -- is the YAML at FLEXKV_RADIXSHMEM_CONFIG_PATH + # (flexkv.common.radixshmem_config; reference docs/radixshmem/config_zh.md). + enable_radixshmem=bool(int(os.getenv('FLEXKV_ENABLE_RADIXSHMEM', 0))), + radixshmem_config_path=os.getenv('FLEXKV_RADIXSHMEM_CONFIG_PATH', '') or None, + # embedded: the bootstrap DP process launches the radix-server subprocess; + # external: a radix-server started by the operator is attached to. + radix_server_launch_mode=os.getenv('FLEXKV_RADIX_SERVER_LAUNCH_MODE', 'embedded').lower(), + # Per-node overrides of the global YAML, for several nodes on one host + # (tests): the node's etcd identity and the bootstrap IP peers dial. Unset + # in a real deployment, where both derive from cluster.rpc_interface. + radix_node_name=os.getenv('FLEXKV_RADIX_NODE_NAME', ''), + radix_rpc_address=os.getenv('FLEXKV_RADIX_RPC_ADDRESS', ''), + index_accel=bool(int(os.getenv('FLEXKV_INDEX_ACCEL', 1))), cpu_layout_type=KVCacheLayoutType(os.getenv('FLEXKV_CPU_LAYOUT', 'BLOCKFIRST').upper()), ssd_layout_type=KVCacheLayoutType(os.getenv('FLEXKV_SSD_LAYOUT', 'BLOCKFIRST').upper()), @@ -1174,6 +1236,7 @@ def update_default_config_from_user_config(rank_info: RankInfo, # Store original GB values for deferred recomputation (when layer_groups become known) cache_config._user_cpu_cache_gb = user_config.cpu_cache_gb cache_config._user_ssd_cache_gb = user_config.ssd_cache_gb + cache_config._num_layers_per_pp_stage = int(rank_info.num_layers_per_pp_stage) cache_config.num_cpu_blocks = ( convert_to_block_num(user_config.cpu_cache_gb, block_size_in_bytes) diff --git a/flexkv/common/radixshmem_config.py b/flexkv/common/radixshmem_config.py new file mode 100644 index 000000000..900b57edf --- /dev/null +++ b/flexkv/common/radixshmem_config.py @@ -0,0 +1,365 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +"""The radixshmem-mode configuration file (``FLEXKV_RADIXSHMEM_CONFIG_PATH``). + +One YAML, identical on every node of a cluster, with five sections: + + cluster / data / index / server + Passed through by key to ``shmradix.ClusterConfig`` / + ``DataPlaneConfig`` / ``IndexConfig`` / ``RadixServerConfig``. Keys are + validated against the dataclass fields of the installed shmradix, so a + new radixshmem field is configurable without a FlexKV change and a typo + fails at startup. Geometry fields (slot counts, slot bytes, alignment, + shm names) are derived from ``CacheConfig`` and rejected here. + client + FlexKV's own RadixClient / prefetch settings. + +Per-node values do not belong in a global file: ``cluster.node_name`` and +``cluster.rpc_address`` are rejected. A node derives its identity from the IP +``cluster.rpc_interface`` resolves to; the two environment variables +``FLEXKV_RADIX_NODE_NAME`` / ``FLEXKV_RADIX_RPC_ADDRESS`` override that for +several nodes on one host (tests). + +``cluster.cluster_id`` is the only namespace: the etcd key prefix and, through +:meth:`RadixShmemConfig.local_id`, every shm / socket / IPC name on this host. + +Reference: ``docs/radixshmem/config_zh.md``. +""" +from __future__ import annotations + +import dataclasses +import threading +from typing import Any, Dict, Optional, Set, Tuple + +import yaml + +from flexkv.common.config import GLOBAL_CONFIG_FROM_ENV + +SECTIONS = ("cluster", "data", "index", "server", "client") + +# Values FlexKV sets differently from radixshmem's own defaults; anything not +# listed takes the shmradix dataclass default. One more default depends on the +# geometry and is resolved in shm_radix_bootstrap.build_radix_server_config: +# index.register_chunk_size = REGISTER_CHUNK_TOKENS // tokens_per_block, so an +# RHT registration chunk covers REGISTER_CHUNK_TOKENS tokens whatever the block +# size (radixshmem's own default is 128 blocks). +REGISTER_CHUNK_TOKENS = 4096 + +FLEXKV_DEFAULTS: Dict[str, Dict[str, Any]] = { + "cluster": { + "cluster_id": "flexkv", + "bootstrap_timeout_sec": 120, + # 1 is a blind overwrite that loses routing entries. + "rht_slots_per_bucket": 4, + }, + "data": {}, + "index": {"data_pool_ratio": 8.0}, + "server": {}, +} + +# Derived from CacheConfig / ModelConfig (shm_radix_bootstrap.expected_geometry) +# or per node; rejected in the file. +FORBIDDEN_KEYS: Dict[str, Set[str]] = { + "cluster": {"node_name", "rpc_address"}, + "data": {"data_bytes", "full_slot_bytes", "swa_slot_bytes", "mamba_slot_bytes", + "slot_align", "data_name"}, + "index": {"name", "tokens_per_block", "full_slots", "swa_slots", "swa_window_blocks", + "mamba_slots", "evict_policy"}, + "server": set(), +} + +_RDMA_TRANSPORTS = {"xrc", "dc"} +_REMOTE_OP_TRANSPORTS = {"zmq", "dc"} +_RHT_SLOTS = {1, 2, 4, 8} + + +class RadixShmemConfigError(ValueError): + """The file is not a valid radixshmem-mode configuration.""" + + +@dataclasses.dataclass(frozen=True) +class RadixClientSettings: + """FlexKV-side settings of the RadixClient and the prefetch path.""" + # Server-side deadline of one prefetch pull; the job completes with the + # local hit when it expires. + prefetch_timeout_ms: int = 5000 + # Peer pulls in flight per CE process before new prefetches skip the peer + # walk; kept under max_outstanding so pull_async never blocks. + prefetch_max_inflight: int = 128 + # Uncollected jobs one RadixClient may hold. + max_outstanding: int = 256 + + +@dataclasses.dataclass(frozen=True) +class RadixShmemConfig: + path: Optional[str] + cluster: Dict[str, Any] + data: Dict[str, Any] + index: Dict[str, Any] + server: Dict[str, Any] + client: RadixClientSettings = RadixClientSettings() + + # ----------------------------------------------------------- cluster + @property + def cluster_id(self) -> str: + return str(self.cluster["cluster_id"]) + + @property + def node_name(self) -> str: + return str(self.cluster.get("node_name", "")) + + @property + def rpc_address(self) -> str: + return str(self.cluster.get("rpc_address", "")) + + @property + def expected_min_nodes(self) -> int: + return int(self.cluster.get("expected_min_nodes", 0)) + + @property + def num_rht_shards(self) -> int: + return int(self.cluster.get("num_rht_shards", 0)) + + @property + def distributed(self) -> bool: + """radixshmem's own criterion (ClusterConfig.distributed).""" + return self.expected_min_nodes > 1 or self.num_rht_shards > 1 + + @property + def bootstrap_timeout_sec(self) -> float: + return float(self.cluster.get("bootstrap_timeout_sec", 60)) + + @property + def attach_timeout_s(self) -> float: + """How long a FlexKV process waits for the radix-server: the cluster + rendezvous plus a margin for SlotStore creation / prefault.""" + return self.bootstrap_timeout_sec + 60.0 + + @property + def local_id(self) -> str: + """Prefix of every name on this host: the index / SlotStore shm, the + gRPC socket, FlexKV's TE channels and GPU registration port. The + cluster id, suffixed with the node name when one was given so that + co-located nodes do not share regions.""" + return f"{self.cluster_id}_{self.node_name}" if self.node_name else self.cluster_id + + # ------------------------------------------------------------ server + @property + def endpoint(self) -> str: + """gRPC endpoint; "" = radixshmem's unix:///dev/shm/.sock.""" + return str(self.server.get("endpoint", "")) + + @property + def hugepage_path(self) -> str: + return str(self.server.get("hugepage_path", "")) + + # ------------------------------------------------------------- tests + def replace_cluster(self, **changes: Any) -> "RadixShmemConfig": + """A copy with ``cluster`` keys changed (test helper; bypasses the + forbidden-key check so node_name / rpc_address can be set).""" + return dataclasses.replace(self, cluster={**self.cluster, **changes}) + + def replace_server(self, **changes: Any) -> "RadixShmemConfig": + return dataclasses.replace(self, server={**self.server, **changes}) + + def describe(self) -> str: + where = self.path or "(defaults)" + s = f"{where}: cluster_id={self.cluster_id}" + if self.distributed: + s += (f", expected_min_nodes={self.expected_min_nodes}, " + f"registry={self.cluster.get('registry')}, " + f"rpc_interface={self.cluster.get('rpc_interface') or '-'}, " + f"rpc_address={self.rpc_address or '-'}, node_name={self.node_name or '(auto)'}") + return s + + +# ------------------------------------------------------------------ loading + +def _shmradix_dataclasses(): + try: + import shmradix + except ImportError as exc: # pragma: no cover + raise ImportError( + "shmradix is not installed; install it from the radixshmem repo " + "(pip install -e radixshmem/python)") from exc + try: + return { + "cluster": shmradix.ClusterConfig, + "data": shmradix.DataPlaneConfig, + "index": shmradix.IndexConfig, + "server": shmradix.RadixServerConfig, + } + except AttributeError as exc: + raise ImportError( + "shmradix lacks the RadixServer configuration dataclasses: FlexKV needs " + "the RadixServer / RadixClient surface of radixshmem") from exc + + +def _read_yaml(path: str) -> Dict[str, Any]: + with open(path) as f: + loaded = yaml.safe_load(f) + if loaded is None: + return {} + if not isinstance(loaded, dict): + raise RadixShmemConfigError(f"{path}: top level must be a mapping of sections") + return loaded + + +def _section(raw: Dict[str, Any], name: str, path: str) -> Dict[str, Any]: + sec = raw.get(name) + if sec is None: + return {} + if not isinstance(sec, dict): + raise RadixShmemConfigError(f"{path}: section '{name}' must be a mapping") + return dict(sec) + + +def _as_list(value: Any) -> Any: + if isinstance(value, str): + return [v.strip() for v in value.split(",") if v.strip()] + return value + + +def _passthrough_section(name: str, given: Dict[str, Any], dc, path: str) -> Dict[str, Any]: + """FlexKV defaults overlaid with the file's keys, validated against the + shmradix dataclass ``dc``.""" + fields = {f.name for f in dataclasses.fields(dc)} + if name == "server": + # RadixServerConfig's nested sections are configured by their own + # sections here, not inline. + fields -= {"index", "data", "cluster"} + forbidden = FORBIDDEN_KEYS[name] & set(given) + if forbidden: + raise RadixShmemConfigError( + f"{path}: '{name}.{sorted(forbidden)[0]}' is not configurable: " + + ("geometry is derived from the FlexKV cache configuration" + if name in ("data", "index") else + "it is a per-node value; set FLEXKV_RADIX_NODE_NAME / " + "FLEXKV_RADIX_RPC_ADDRESS on that node instead")) + unknown = set(given) - fields + if unknown: + raise RadixShmemConfigError( + f"{path}: unknown key(s) in '{name}': {sorted(unknown)}; " + f"shmradix.{dc.__name__} has {sorted(fields)}") + merged = {**FLEXKV_DEFAULTS[name], **given} + if "transfer_devices" in merged: + merged["transfer_devices"] = [str(d) for d in _as_list(merged["transfer_devices"])] + if "rht_shard_holders" in merged: + merged["rht_shard_holders"] = [int(r) for r in _as_list(merged["rht_shard_holders"])] + return merged + + +def _client_section(given: Dict[str, Any], path: str) -> RadixClientSettings: + fields = {f.name for f in dataclasses.fields(RadixClientSettings)} + unknown = set(given) - fields + if unknown: + raise RadixShmemConfigError( + f"{path}: unknown key(s) in 'client': {sorted(unknown)}; expected {sorted(fields)}") + return RadixClientSettings(**{k: int(v) for k, v in given.items()}) + + +def _validate(cfg: RadixShmemConfig, path: str) -> None: + c = cfg.cluster + if not cfg.cluster_id: + raise RadixShmemConfigError(f"{path}: cluster.cluster_id must not be empty") + if cfg.distributed: + if not c.get("registry"): + raise RadixShmemConfigError( + f"{path}: cluster mode (expected_min_nodes > 1) needs cluster.registry, " + f"e.g. 'etcd://10.0.0.1:2379'") + if not c.get("rpc_interface") and not cfg.rpc_address: + raise RadixShmemConfigError( + f"{path}: cluster mode needs cluster.rpc_interface (the NIC whose IP peers " + f"dial and this node's identity derives from) or FLEXKV_RADIX_RPC_ADDRESS") + if cfg.rpc_address == "0.0.0.0": + raise RadixShmemConfigError( + "FLEXKV_RADIX_RPC_ADDRESS=0.0.0.0 gives every node the same identity; " + "use this node's address") + if cfg.expected_min_nodes > 0 and cfg.num_rht_shards > cfg.expected_min_nodes: + raise RadixShmemConfigError( + f"{path}: cluster.num_rht_shards={cfg.num_rht_shards} exceeds " + f"expected_min_nodes={cfg.expected_min_nodes}; there cannot be more RHT shard " + f"holders than nodes") + slots = int(c.get("rht_slots_per_bucket", 1)) + if slots not in _RHT_SLOTS: + raise RadixShmemConfigError( + f"{path}: cluster.rht_slots_per_bucket={slots} must be one of {sorted(_RHT_SLOTS)}") + for key in ("rht_transport", "peer_index_transport"): + val = c.get(key) + if val is not None and val not in _RDMA_TRANSPORTS: + raise RadixShmemConfigError( + f"{path}: cluster.{key}={val!r} must be one of {sorted(_RDMA_TRANSPORTS)}") + rot = c.get("remote_op_transport") + if rot is not None and rot not in _REMOTE_OP_TRANSPORTS: + raise RadixShmemConfigError( + f"{path}: cluster.remote_op_transport={rot!r} must be one of " + f"{sorted(_REMOTE_OP_TRANSPORTS)}") + if cfg.client.prefetch_max_inflight >= cfg.client.max_outstanding: + raise RadixShmemConfigError( + f"{path}: client.prefetch_max_inflight={cfg.client.prefetch_max_inflight} must be " + f"below client.max_outstanding={cfg.client.max_outstanding}, or pull_async blocks") + if cfg.client.prefetch_timeout_ms <= 0: + raise RadixShmemConfigError(f"{path}: client.prefetch_timeout_ms must be > 0") + + +def load_radixshmem_config(path: Optional[str] = None, + *, + node_name: str = "", + rpc_address: str = "") -> RadixShmemConfig: + """Parse ``path`` (None or "" = all defaults, i.e. standalone) and apply + the per-node overrides. Raises :class:`RadixShmemConfigError` on an + invalid file, ``ImportError`` without shmradix.""" + dcs = _shmradix_dataclasses() + label = path or "(defaults)" + raw = _read_yaml(path) if path else {} + unknown = set(raw) - set(SECTIONS) + if unknown: + raise RadixShmemConfigError( + f"{label}: unknown section(s) {sorted(unknown)}; expected {list(SECTIONS)}") + sections = {name: _passthrough_section(name, _section(raw, name, label), dcs[name], label) + for name in ("cluster", "data", "index", "server")} + if node_name: + sections["cluster"]["node_name"] = str(node_name) + if rpc_address: + sections["cluster"]["rpc_address"] = str(rpc_address) + # radixshmem lets the interface win over the address; an explicit + # per-node address means the global interface must not apply here. + sections["cluster"]["rpc_interface"] = "" + cfg = RadixShmemConfig(path=path or None, client=_client_section( + _section(raw, "client", label), label), **sections) + _validate(cfg, label) + return cfg + + +# -------------------------------------------------------------- singleton + +_lock = threading.Lock() +_cached: Optional[Tuple[Tuple[str, str, str], RadixShmemConfig]] = None + + +def _env_key() -> Tuple[str, str, str]: + env = GLOBAL_CONFIG_FROM_ENV + return (str(env.radixshmem_config_path or ""), str(env.radix_node_name or ""), + str(env.radix_rpc_address or "")) + + +def get_radixshmem_config() -> RadixShmemConfig: + """The process's configuration: loaded from ``GLOBAL_CONFIG_FROM_ENV`` + (``FLEXKV_RADIXSHMEM_CONFIG_PATH`` + the two per-node overrides) on first + use and whenever those three values change.""" + global _cached + key = _env_key() + with _lock: + if _cached is None or _cached[0] != key: + path, node_name, rpc_address = key + _cached = (key, load_radixshmem_config(path or None, node_name=node_name, + rpc_address=rpc_address)) + return _cached[1] + + +def set_radixshmem_config(cfg: Optional[RadixShmemConfig]) -> None: + """Install ``cfg`` as the process's configuration (tests); None reverts to + loading from the environment.""" + global _cached + with _lock: + _cached = None if cfg is None else (_env_key(), cfg) diff --git a/flexkv/integration/config.py b/flexkv/integration/config.py index 813048ddc..427ee8b1c 100644 --- a/flexkv/integration/config.py +++ b/flexkv/integration/config.py @@ -21,6 +21,56 @@ def _dsv4_swa_transfer_enabled_from_env() -> bool: return bool(int(os.getenv("FLEXKV_ENABLE_SWA_TRANSFER", "1"))) +def _resolve_vllm_dp_rank(parallel_config: object) -> int: + """This engine's DP rank, as FlexKV needs it rather than as vLLM reports it. + + For non-MoE models vLLM runs each DP rank as a fully independent engine and + resets that child's ``data_parallel_rank`` to 0, keeping the real rank only + in ``data_parallel_index`` (vllm/v1/engine/core.py, "Non-MoE DP ranks are + completely independent, so treat like DP=1"). FlexKV cannot follow suit: its + per-engine identity ``dp_client_id = instance_id * dp_size + dp_rank`` is + what keeps the DP engines of one node apart, so a collapsed rank makes all + of them claim shm-radix bootstrap ownership, transfer-engine channel 0, the + same gpu_register endpoint, and overlapping graph/op id ranges. + + ``FLEXKV_DP_RANK`` overrides both, for launchers that pin the rank + themselves. + """ + env_rank = os.environ.get("FLEXKV_DP_RANK") + if env_rank: + return int(env_rank) + dp_index = getattr(parallel_config, "data_parallel_index", None) + if dp_index is not None: + return int(dp_index) + return int(getattr(parallel_config, "data_parallel_rank", 0)) + + +def _resolve_vllm_dp_size(parallel_config: object) -> int: + """The node's DP width. Needs an env var: vLLM erases it in the children. + + The same non-MoE branch that collapses the rank also sets the child's + ``data_parallel_size`` to 1, and unlike the rank it leaves no surviving copy + anywhere in ``parallel_config``. FlexKV derives ``total_clients`` (the shared + transfer engine's channel count and its expected GPU-registration count) and + ``total_gpus`` (which gates clearing ``CUDA_VISIBLE_DEVICES`` in the TE + subprocess so it can open every DP rank's IPC handles) from it, so a stale 1 + leaves the TE waiting for registrations that already arrived under a + duplicate device id. + + Only ever raises the value, so single-engine and MoE setups -- where vLLM's + own number is already right -- are untouched. + """ + dp_size = int(getattr(parallel_config, "data_parallel_size", 1)) + env_size = os.environ.get("FLEXKV_DP_SIZE") + if env_size and int(env_size) > dp_size: + logger.info( + f"[FlexKV vllm] dp_size {dp_size} -> {env_size} from FLEXKV_DP_SIZE " + f"(vLLM resets data_parallel_size to 1 in non-MoE DP children)" + ) + return int(env_size) + return dp_size + + def _is_nvfp4_dtype_str(dtype_str: Optional[str]) -> bool: """Return True if *dtype_str* selects the NVFP4 packed KV cache layout.""" return isinstance(dtype_str, str) and dtype_str.lower() in ("nvfp4", "fp4", "e2m1") @@ -109,6 +159,32 @@ def __post_init__(self): if self.gpu_register_port == "": self.gpu_register_port = self.server_recv_port + "_gpu_register" + @staticmethod + def get_sglang_node_local_dp_size( + server_args: object, + ) -> Optional[int]: + """Return a safe node-local DP width for SGLang DP Attention. + + SGLang validates the composite TP/CP dimensions and assigns contiguous + DP groups. With ``pp_size == 1``, groups are node-local exactly when + ``dp_size`` is evenly divisible by ``nnodes``. + + ``None`` keeps the existing cross-node TP/PP topology unchanged. + """ + dp_size = max(1, int(getattr(server_args, "dp_size", 1) or 1)) + pp_size = max(1, int(getattr(server_args, "pp_size", 1))) + nnodes = max(1, int(getattr(server_args, "nnodes", 1))) + + if ( + not bool(getattr(server_args, "enable_dp_attention", False)) + or dp_size == 1 + or nnodes == 1 + or pp_size != 1 + or dp_size % nnodes != 0 + ): + return None + return dp_size // nnodes + def _resolve_dtype( self, framework_dtype_str: Optional[str], @@ -201,7 +277,7 @@ def post_init_from_vllm_config( parallel_config = vllm_config.parallel_config tp_rank = int(getattr(parallel_config, 'tensor_parallel_rank', 0)) pp_rank = int(getattr(parallel_config, 'pipeline_parallel_rank', 0)) - dp_rank = int(getattr(parallel_config, 'data_parallel_rank', 0)) + dp_rank = _resolve_vllm_dp_rank(parallel_config) node_rank = int(getattr(parallel_config, 'node_rank', 0)) self.cache_config.tokens_per_block = vllm_config.cache_config.block_size @@ -214,7 +290,7 @@ def post_init_from_vllm_config( # kv_dim/num_kv_heads are derived from the physical tensor shape # (cache_shape ndim), not from framework MLA flags. self.model_config.tp_size = int(parallel_config.tensor_parallel_size) - self.model_config.dp_size = int(parallel_config.data_parallel_size) + self.model_config.dp_size = _resolve_vllm_dp_size(parallel_config) self.model_config.pp_size = int(parallel_config.pipeline_parallel_size) # vLLM CP (context parallel) support: read cp_size from parallel_config. # Falls back to 1 if the attribute is not present (older vLLM versions). @@ -520,9 +596,36 @@ def post_init_from_sglang_config( enable_dp_attention = bool(server_args.enable_dp_attention) attn_cp_size = int(getattr(server_args, 'attn_cp_size', 1)) kv_cache_dtype = getattr(server_args, 'kv_cache_dtype', None) - + # Node-local DP is a property of the SGLang placement, not of the tier + # underneath it: with DP attention and pp_size == 1, a dp_size divisible + # by nnodes puts every DP group on a single node, so FlexKV forms one + # instance per node whatever the shared tier is (radix-shmem, or + # mooncake-store, where the store itself carries cross-node reuse). + # get_sglang_node_local_dp_size returns None for every placement where + # that does not hold, which leaves the cross-node TP/PP path untouched. + local_dp_size = self.get_sglang_node_local_dp_size(server_args) + + if dp_rank is None and GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and sglang_dp_size > 1: + # Every DP process would derive dp_client_id 0: the same radix-server + # bootstrap ownership, TE channel and graph/op id range. + raise ValueError( + "[FlexKV SGLang] radix_shmem with dp_size > 1 needs the scheduler's " + "dp_rank; got None") dp_rank = 0 if dp_rank is None else int(dp_rank) cp_rank = 0 if cp_rank is None else int(cp_rank) + if local_dp_size is not None: + logger.info( + "[FlexKV SGLang] Enabling node-local DP (one FlexKV instance per " + "node): global_dp_size=%d, local_dp_size=%d, node_rank=%d", + sglang_dp_size, + local_dp_size, + int(node_rank), + ) + elif enable_dp_attention and sglang_dp_size > 1 and int(nnodes) > 1: + logger.warning( + "[FlexKV SGLang] Node-local DP is not available for this DP/PP " + "placement; preserving the legacy cross-node path." + ) attn_dp_size = sglang_dp_size if enable_dp_attention else 1 attn_tp_size = max(1, sglang_tp_size // (attn_dp_size * attn_cp_size)) @@ -637,6 +740,7 @@ def post_init_from_sglang_config( pp_end_layer = self.model_config.num_layers self.model_config.enable_dp_attention = bool(enable_dp_attention) self.model_config.nnodes = max(1, int(nnodes)) + self.model_config.local_dp_size = local_dp_size _dist_init_addr = getattr(server_args, 'dist_init_addr', None) if _dist_init_addr and int(nnodes) > 1: self.model_config.master_host = _dist_init_addr.split(":")[0] @@ -703,6 +807,9 @@ def post_init_from_sglang_config( enabled=True, num_swa_layers=self.model_config.num_layers, bytes_per_token_per_layer=swa_bytes_per_token, + # DSv4's 128-token window fits inside one 256-token page, + # so a window is a single slot (radixshmem W=1). + window_blocks=1, ) # Gate the SWA data plane (byte movement) behind an env switch so it # can be turned off for A/B or if a byte-layout issue surfaces in diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py new file mode 100644 index 000000000..0a27c850d --- /dev/null +++ b/flexkv/server/shm_radix_bootstrap.py @@ -0,0 +1,483 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +Bootstrap for the radixshmem-backed CPU tier. + +One ``radix-server`` process per node owns the radix index shm, the SlotStore +(the CPU KV pool: one slot per block, a FULL pool and optionally an SWA pool), +the RDMA transfer engine and the etcd data-plane entry. In this mode FlexKV +neither allocates CPU KV memory nor moves bytes between nodes itself: + + * the bootstrap DP process (instance 0, dp 0) launches the server from the + FlexKV configuration -- geometry from ``CacheConfig``, everything else from + the YAML at ``FLEXKV_RADIXSHMEM_CONFIG_PATH`` (``flexkv.common.radixshmem_config``) + -- (``FLEXKV_RADIX_SERVER_LAUNCH_MODE=embedded``) or expects one started by + the operator (``external``); + * every DP scheduler process, the TE process and its transfer workers attach + with ``shmradix.RadixClient(name)``: index operations, ``store`` (the + SlotStore mapping) and ``pull_async`` (the server-side peer pull). + +Naming: index ``/shmradix__cpu`` where ``local_id`` is the YAML's +``cluster.cluster_id`` (plus ``_`` when FLEXKV_RADIX_NODE_NAME names +one of several co-located nodes). A cluster node's index gets ``_`` +appended by radixshmem itself and is resolved through the gRPC socket +``/dev/shm/shmradix__cpu.sock``, so attachers only need the base +name. The SlotStore is ``_data``. + +Geometry: FlexKV stays the source of slot counts and slot bytes. One FULL slot +holds exactly one CPU block as ``StorageEngine`` lays it out (BLOCKFIRST: all +layers of a block contiguous), one SWA slot one SWA page. ``slot_align`` is +chosen so that the SlotStore stride equals the block size exactly, which lets +the H2D / D2H workers address the pool with the strides of a plain tensor. The +TE re-checks the attached regions against the layouts it builds (``check_geometry``). +""" +from __future__ import annotations + +import dataclasses +import multiprocessing as mp +import os +import signal +import time +from typing import Any, Dict, List, Optional + +import torch + +from flexkv.common.config import (GLOBAL_CONFIG_FROM_ENV, CacheConfig, LayerGroupSpec, + ModelConfig, SWAPoolConfig) +from flexkv.common.debug import flexkv_logger +from flexkv.common.radixshmem_config import (REGISTER_CHUNK_TOKENS, RadixShmemConfig, + get_radixshmem_config) +from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType + +try: + import shmradix +except ImportError: # pragma: no cover + shmradix = None + + +_SHM_PREFIX = "/shmradix" +# Pool bases are page aligned regardless; a larger per-slot alignment only pads. +_MAX_SLOT_ALIGN = 4096 +DEFAULT_HUGETLBFS_DIR = "/mnt/hugepages" + + +def _ensure_shmradix() -> None: + if shmradix is None: + raise ImportError( + "shmradix is not installed; install it from the radixshmem repo " + "(pip install -e radixshmem/python)") + for name in ("RadixServer", "RadixServerConfig", "IndexConfig", "DataPlaneConfig", + "ClusterConfig", "RadixClient"): + if not hasattr(shmradix, name): + raise ImportError( + f"shmradix lacks {name}: FlexKV needs the RadixServer / RadixClient " + f"surface of radixshmem (transfer-server branch or later)") + + +def radix_index_name(local_id: str) -> str: + """Base shm name of the CPU tier's index for ``local_id`` + (``RadixShmemConfig.local_id``); it also names the gRPC socket.""" + return f"{_SHM_PREFIX}_{local_id}_cpu" + + +def radix_data_name(local_id: str) -> str: + return radix_index_name(local_id) + "_data" + + +def radix_socket_path(local_id: str) -> str: + return "/dev/shm/" + radix_index_name(local_id).lstrip("/").replace("/", "_") + ".sock" + + +# ------------------------------------------------------------------ geometry + +def _resolve_groups(groups: Optional[List[LayerGroupSpec]], + default_dtype: torch.dtype) -> Optional[List[LayerGroupSpec]]: + if groups is None: + return None + return [g if g.dtype is not None else dataclasses.replace(g, dtype=default_dtype) + for g in groups] + + +def num_layers_per_pp_stage(model_config: ModelConfig, cache_config: CacheConfig) -> int: + """Layers one CPU block covers: what the adapter recorded, else an even split.""" + recorded = int(getattr(cache_config, "_num_layers_per_pp_stage", 0) or 0) + if recorded > 0: + return recorded + return max(1, model_config.num_layers // max(1, model_config.pp_size)) + + +def layout_block_bytes(layout: KVCacheLayout, dtype: torch.dtype) -> int: + """Bytes of one block of ``layout`` (a multi-group layout is byte-flat).""" + if layout.layer_groups is not None: + return int(layout.get_block_stride()) + return int(layout.get_block_stride()) * dtype.itemsize + + +def cpu_kv_layout(model_config: ModelConfig, cache_config: CacheConfig, + num_blocks: int) -> KVCacheLayout: + """The CPU FULL layout exactly as ``StorageEngine`` builds it.""" + return KVCacheLayout( + type=GLOBAL_CONFIG_FROM_ENV.cpu_layout_type, + num_layer=num_layers_per_pp_stage(model_config, cache_config), + num_block=num_blocks, + tokens_per_block=cache_config.tokens_per_block, + num_head=model_config.num_kv_heads_per_node, + head_size=model_config.head_size, + kv_dim=model_config.kv_dim, + num_kv_heads=model_config.num_kv_heads, + layer_groups=_resolve_groups(model_config.layer_groups, model_config.dtype), + tp_size=model_config.tp_size, + ) + + +def cpu_block_bytes(model_config: ModelConfig, cache_config: CacheConfig) -> int: + return layout_block_bytes(cpu_kv_layout(model_config, cache_config, 1), model_config.dtype) + + +def swa_pool_config(cache_config: CacheConfig) -> Optional[SWAPoolConfig]: + swa = cache_config.swa + if swa is None or not swa.enabled or swa.num_slots <= 0: + return None + return swa + + +def swa_cpu_kv_layout(model_config: ModelConfig, cache_config: CacheConfig, + num_blocks: int) -> KVCacheLayout: + """The CPU SWA layout exactly as ``StorageEngine`` builds it (uint8, one page + per slot; DSv4 sidecar groups come from ``SWAPoolConfig.layer_groups``).""" + swa = cache_config.swa + return KVCacheLayout( + type=GLOBAL_CONFIG_FROM_ENV.cpu_layout_type, + num_layer=swa.num_swa_layers, + num_block=num_blocks, + tokens_per_block=cache_config.tokens_per_block, + num_head=1, + head_size=swa.bytes_per_token_per_layer, + kv_dim=1, + num_kv_heads=1, + layer_groups=_resolve_groups(swa.layer_groups, torch.uint8), + tp_size=model_config.tp_size, + ) + + +def swa_block_bytes(model_config: ModelConfig, cache_config: CacheConfig) -> int: + return layout_block_bytes(swa_cpu_kv_layout(model_config, cache_config, 1), torch.uint8) + + +def slot_align_for(*sizes: int) -> int: + """Largest power of two <= 4096 dividing every size. radixshmem rounds each + slot's stride up to ``slot_align``, so this keeps stride == slot bytes.""" + align = _MAX_SLOT_ALIGN + for size in sizes: + if size <= 0: + continue + while size % align: + align //= 2 + return max(align, 1) + + +@dataclasses.dataclass(frozen=True) +class RadixGeometry: + """What FlexKV expects the server's regions to look like.""" + tokens_per_block: int + full_slots: int + full_slot_bytes: int + swa_slots: int = 0 + swa_slot_bytes: int = 0 + swa_window_blocks: int = 0 + + @property + def slot_align(self) -> int: + return slot_align_for(self.full_slot_bytes, self.swa_slot_bytes) + + @property + def data_bytes(self) -> int: + return self.full_slots * self.full_slot_bytes + self.swa_slots * self.swa_slot_bytes + + def describe(self) -> str: + s = (f"tokens_per_block={self.tokens_per_block}, FULL {self.full_slots} x " + f"{self.full_slot_bytes} B") + if self.swa_slots: + s += (f", SWA {self.swa_slots} x {self.swa_slot_bytes} B " + f"(window {self.swa_window_blocks})") + return s + f", slot_align={self.slot_align}, data={self.data_bytes / 2**30:.2f} GiB" + + +def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> RadixGeometry: + if GLOBAL_CONFIG_FROM_ENV.cpu_layout_type != KVCacheLayoutType.BLOCKFIRST: + raise ValueError( + "radixshmem needs FLEXKV_CPU_LAYOUT=BLOCKFIRST: one SlotStore slot is one " + "contiguous block, which LAYERFIRST does not give") + if cache_config.num_cpu_blocks <= 0: + raise ValueError(f"cache_config.num_cpu_blocks={cache_config.num_cpu_blocks} must be > 0") + geo = RadixGeometry( + tokens_per_block=cache_config.tokens_per_block, + full_slots=int(cache_config.num_cpu_blocks), + full_slot_bytes=cpu_block_bytes(model_config, cache_config), + ) + swa = swa_pool_config(cache_config) + if swa is not None: + if swa.window_blocks < 1: + raise ValueError( + f"cache_config.swa.window_blocks={swa.window_blocks} must be >= 1") + if swa.num_slots < swa.window_blocks: + # All-or-none window allocation: a pool smaller than one window can + # never store anything, so fail at startup. + raise ValueError( + f"cache_config.swa.num_slots={swa.num_slots} cannot hold one " + f"{swa.window_blocks}-block SWA window; raise num_slots or disable SWA") + geo = dataclasses.replace( + geo, swa_slots=int(swa.num_slots), + swa_slot_bytes=swa_block_bytes(model_config, cache_config), + swa_window_blocks=int(swa.window_blocks)) + return geo + + +def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, + label: str = "radixshmem") -> None: + """Fail closed when the attached regions differ from FlexKV's own layout: a + slot count or stride mismatch would otherwise become a silent misaddressed + transfer.""" + g = client.geometry + pools = g["pools"] + diffs: List[str] = [] + if int(g["block_size"]) != expected.tokens_per_block: + diffs.append(f"tokens_per_block server={g['block_size']} flexkv={expected.tokens_per_block}") + full = pools["full"] + if int(full["num_slots"]) != expected.full_slots: + diffs.append(f"FULL slots server={full['num_slots']} flexkv={expected.full_slots}") + if int(full["slot_bytes"]) != expected.full_slot_bytes: + diffs.append(f"FULL slot_bytes server={full['slot_bytes']} flexkv={expected.full_slot_bytes}") + if not client.info.data_plane: + diffs.append("server is index-only (no SlotStore); FlexKV needs the data plane") + else: + store = client.store + stride = int(store.pool(shmradix.ComponentType.FULL).slot_bytes) + if stride != expected.full_slot_bytes: + diffs.append(f"FULL stride server={stride} flexkv={expected.full_slot_bytes} " + f"(slot_align must divide the block size)") + swa = pools.get("swa") + if expected.swa_slots > 0: + if swa is None: + diffs.append("server has no SWA pool but FlexKV's SWA tier is on") + else: + if int(swa["num_slots"]) != expected.swa_slots: + diffs.append(f"SWA slots server={swa['num_slots']} flexkv={expected.swa_slots}") + if int(swa["slot_bytes"]) != expected.swa_slot_bytes: + diffs.append(f"SWA slot_bytes server={swa['slot_bytes']} flexkv={expected.swa_slot_bytes}") + if int(swa.get("window_blocks", 0)) != expected.swa_window_blocks: + diffs.append(f"SWA window server={swa.get('window_blocks')} " + f"flexkv={expected.swa_window_blocks}") + if client.info.data_plane: + stride = int(client.store.pool(shmradix.ComponentType.SWA).slot_bytes) + if stride != expected.swa_slot_bytes: + diffs.append(f"SWA stride server={stride} flexkv={expected.swa_slot_bytes}") + elif swa is not None: + diffs.append("server has an SWA pool that FlexKV's configuration does not") + if diffs: + raise ValueError( + f"{label}: the attached radixshmem regions do not match FlexKV's configuration " + f"({expected.describe()}): " + "; ".join(diffs)) + + +# ------------------------------------------------------------- server config + +def build_radix_server_config(model_config: ModelConfig, + cache_config: CacheConfig, + rcfg: Optional[RadixShmemConfig] = None, + ) -> "shmradix.RadixServerConfig": + """The one radix-server this FlexKV node needs: index sized from + ``cache_config`` (FULL slots = ``num_cpu_blocks``, SWA slots = ``swa.num_slots``), + SlotStore sized so that every slot is exactly one block, and the cluster / + data-plane / index / server settings of ``rcfg`` (default: the process's + ``FLEXKV_RADIXSHMEM_CONFIG_PATH``) passed through. Raises on an + inconsistent configuration.""" + _ensure_shmradix() + if rcfg is None: + rcfg = get_radixshmem_config() + geo = expected_geometry(model_config, cache_config) + index_kwargs = dict(rcfg.index) + # an RHT registration chunk covers REGISTER_CHUNK_TOKENS tokens unless the file says otherwise + index_kwargs.setdefault("register_chunk_size", + max(1, REGISTER_CHUNK_TOKENS // geo.tokens_per_block)) + index = shmradix.IndexConfig( + name=radix_index_name(rcfg.local_id), + tokens_per_block=geo.tokens_per_block, + full_slots=geo.full_slots, + swa_slots=geo.swa_slots, + swa_window_blocks=geo.swa_window_blocks, + **index_kwargs, + ) + data = shmradix.DataPlaneConfig( + data_bytes=geo.data_bytes, + full_slot_bytes=geo.full_slot_bytes, + swa_slot_bytes=geo.swa_slot_bytes, + slot_align=geo.slot_align, + data_name=radix_data_name(rcfg.local_id), + **rcfg.data, + ) + cluster = shmradix.ClusterConfig(**rcfg.cluster) + server_kwargs = dict(rcfg.server) + if not server_kwargs.get("hugepage_path") and cache_config.use_hugepage_cpu_buffer: + server_kwargs["hugepage_path"] = os.environ.get("FLEXKV_HUGETLBFS_DIR", + DEFAULT_HUGETLBFS_DIR) + cfg = shmradix.RadixServerConfig(index=index, data=data, cluster=cluster, **server_kwargs) + flexkv_logger.info( + f"radixshmem server config for {index.name}: {geo.describe()}, " + f"{rcfg.describe()}, register_chunk_size={index.register_chunk_size}, " + f"hugepage_path={cfg.hugepage_path or '(shm)'}, " + f"prefault={data.prefault}, transfer_devices={data.transfer_devices or '(all)'}") + return cfg + + +# ------------------------------------------------------------ server process + +def _radix_server_main(cfg, ready, stop, conn) -> None: + """Body of the radix-server subprocess: bring the server up, report, wait.""" + import shmradix as _shmradix + + def _on_term(signum, frame): # noqa: ARG001 + stop.set() + + signal.signal(signal.SIGTERM, _on_term) + signal.signal(signal.SIGINT, _on_term) + try: + server = _shmradix.RadixServer(cfg) + server.start() + except BaseException as e: # noqa: BLE001 - reported to the parent, which raises + try: + conn.send(("error", f"{type(e).__name__}: {e}")) + finally: + conn.close() + return + index = server.index + try: + conn.send(("ready", { + "index_name": index.shm_name(), + "rank": int(index.rank()), + "world_size": int(index.world_size()), + "distributed": bool(index.is_distributed()), + })) + finally: + conn.close() + ready.set() + try: + while not stop.wait(0.5): + pass + except KeyboardInterrupt: + pass + finally: + server.close() + + +class RadixServerProcess: + """The embedded radix-server: a spawned subprocess running ``RadixServer``. + + Not the scheduler process (its gRPC threads and transfer polling thread + would contend for the GIL, and a clustered server holds RDMA contexts that + do not survive a fork) and not the TE process (which cannot come up before + the GPU registrations, while the CEs attach the index at construction). + """ + + def __init__(self, cfg: "shmradix.RadixServerConfig"): + self.cfg = cfg + self._ctx = mp.get_context("spawn") + self._ready = self._ctx.Event() + self._stop = self._ctx.Event() + self.process = None + self.info: Dict[str, Any] = {} + + def start(self, timeout_s: Optional[float] = None) -> "RadixServerProcess": + if timeout_s is None: + timeout_s = float(self.cfg.cluster.bootstrap_timeout_sec) + 60.0 + parent, child = self._ctx.Pipe(duplex=False) + self.process = self._ctx.Process( + target=_radix_server_main, + args=(self.cfg, self._ready, self._stop, child), + name="flexkv-radix-server", + daemon=True, + ) + self.process.start() + child.close() + deadline = time.monotonic() + timeout_s + try: + while True: + if parent.poll(0.2): + kind, payload = parent.recv() + break + if not self.process.is_alive(): + raise RuntimeError("radix-server exited during startup (see its log)") + if time.monotonic() > deadline: + self.shutdown() + raise TimeoutError( + f"radix-server did not become ready within {timeout_s:.0f}s " + f"(cluster rendezvous or SlotStore prefault still pending?)") + finally: + parent.close() + if kind == "error": + self.shutdown() + raise RuntimeError(f"radix-server failed to start: {payload}") + self.info = payload + flexkv_logger.info( + f"radix-server pid={self.process.pid} ready: index={payload['index_name']} " + f"rank={payload['rank']}/{payload['world_size']} distributed={payload['distributed']}") + return self + + @property + def cluster_rank(self) -> int: + return int(self.info.get("rank", 0)) + + def shutdown(self, timeout: float = 15.0) -> None: + if self.process is None: + return + self._stop.set() + self.process.join(timeout) + if self.process.is_alive(): + self.process.terminate() + self.process.join(5.0) + self.process = None + + +# ------------------------------------------------------------------- attach + +def attach_radix_client(name: str, + timeout_s: Optional[float] = None, + *, + rcfg: Optional[RadixShmemConfig] = None, + max_outstanding: Optional[int] = None) -> "shmradix.RadixClient": + """``shmradix.RadixClient(name)``, retried until the server's socket exists + and the server is ready (an embedded server starts concurrently with the + CEs, an external one may still be rendezvousing). Endpoint, timeout and + ``max_outstanding`` default to ``rcfg`` (the process's configuration). + + The returned client owns the index attach, the SlotStore mapping and the + gRPC channel; keep it alive for as long as its slots are addressed. + """ + _ensure_shmradix() + if rcfg is None: + rcfg = get_radixshmem_config() + if timeout_s is None: + timeout_s = rcfg.attach_timeout_s + if max_outstanding is None: + max_outstanding = rcfg.client.max_outstanding + deadline = time.monotonic() + timeout_s + last: Optional[BaseException] = None + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError( + f"radix-server {name} not attachable within {timeout_s:.0f}s: {last}") + try: + return shmradix.RadixClient(name, endpoint=rcfg.endpoint or None, + timeout_s=max(1.0, remaining), + max_outstanding=max_outstanding) + except Exception as e: # noqa: BLE001 - socket not there yet, server starting + last = e + flexkv_logger.debug(f"attach to radix-server {name} failed (will retry): {e}") + time.sleep(0.2) + + +def radix_cluster_rank(client: "shmradix.RadixClient") -> int: + """The cluster rank etcd assigned this node (0 when standalone).""" + rank = int(getattr(client.info, "rank", -1)) + return rank if rank >= 0 else int(client.rank()) From 1a8c0982641b340833456a3b7d568c83d68e2288 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:27:40 +0800 Subject: [PATCH 04/21] radixshmem: CPU tier on a radix-server (index engine, planner, SlotStore data plane) * flexkv/cache/radix_shmem_engine.py: CacheEngineRadixShmem, a RadixClient on the node's shared index; pinned matches (ShmRadixMatch), publish-after- transfer inserts (StagedRadixInsert), FULL and SWA components, peer prefetch via RadixClient.pull_async. * flexkv/cache/radix_shmem_planner.py: RadixShmemCacheEngine, a GlobalCacheEngine subclass that plans GET / PUT / PREFETCH on that tier and hands the prefetch job to KVTaskEngine on a RadixPlanHandle. * flexkv/cache/cache_engine.py: the hooks the subclass needs -- _build_cpu_cache_engine, and the shared GET/PUT prologue _prepare_request returning a RequestWindow. * flexkv/storage/storage_engine.py, flexkv/storage/allocator.py: in radixshmem mode the CPU FULL / SWA pools are views of the radix-server's SlotStore (_attach_radix_pool); workers re-attach them by name through SlotStoreTensorHandle in worker_data. Registered in main's single (PoolEndpoint, device_id) handle registry. Part 4/7 of the radixshmem rebase (see part 1 for provenance). Co-authored-by: Hao Xu Co-authored-by: Iris Ge Co-authored-by: linhu-nv Co-authored-by: teeebin --- flexkv/cache/cache_engine.py | 159 ++++---- flexkv/cache/radix_shmem_engine.py | 362 +++++++++++++++++ flexkv/cache/radix_shmem_planner.py | 580 ++++++++++++++++++++++++++++ flexkv/storage/allocator.py | 53 ++- flexkv/storage/storage_engine.py | 105 ++++- 5 files changed, 1171 insertions(+), 88 deletions(-) create mode 100644 flexkv/cache/radix_shmem_engine.py create mode 100644 flexkv/cache/radix_shmem_planner.py diff --git a/flexkv/cache/cache_engine.py b/flexkv/cache/cache_engine.py index 900ec6063..c74ea5f27 100644 --- a/flexkv/cache/cache_engine.py +++ b/flexkv/cache/cache_engine.py @@ -310,6 +310,18 @@ class SWAReadReservation: h2d_id: int +@dataclass(frozen=True) +class RequestWindow: + """A GET/PUT request reduced to whole blocks: the masked block range + ``[block_start_idx, block_end_idx)``, the GPU blocks it maps to, and the + SequenceMeta of the block-aligned token prefix.""" + + block_start_idx: int + block_end_idx: int + gpu_block_ids: np.ndarray + sequence_meta: SequenceMeta + + class TransferPlanHandle: """Completion callback for a planned get/put, with an abort path. @@ -992,37 +1004,7 @@ def __init__(self, cache_config: CacheConfig, model_config: ModelConfig, redis_m ) if cache_config.enable_cpu: - if cache_config.enable_p2p_cpu: - self.cpu_cache_engine = HierarchyLRCacheEngine.from_cache_config( - cache_config, self.node_id, DeviceType.CPU, meta=self.redis_meta) - elif self.index_accel: - self.cpu_cache_engine = CacheEngineAccel( - device_type=DeviceType.CPU, - num_total_blocks=cache_config.num_cpu_blocks, - tokens_per_block=cache_config.tokens_per_block, - evict_ratio=self.evict_ratio, - hit_reward_seconds=self.hit_reward_seconds, - evict_start_threshold=self.evict_start_threshold, - eviction_policy=self.eviction_policy, - event_collector=event_collector, - metrics_collector=self._metrics_collector, - protected_threshold=self.protected_threshold, - swa_config=cache_config.swa, - ) - else: - self.cpu_cache_engine = CacheEngine( - device_type=DeviceType.CPU, - num_total_blocks=cache_config.num_cpu_blocks, - tokens_per_block=cache_config.tokens_per_block, - evict_ratio=self.evict_ratio, - hit_reward_seconds=self.hit_reward_seconds, - evict_start_threshold=self.evict_start_threshold, - eviction_policy=self.eviction_policy, - event_collector=event_collector, - metrics_collector=self._metrics_collector, - protected_threshold=self.protected_threshold, - swa_config=cache_config.swa, - ) + self.cpu_cache_engine = self._build_cpu_cache_engine(cache_config, event_collector) self.cache_engines[DeviceType.CPU] = self.cpu_cache_engine if cache_config.enable_ssd: if cache_config.enable_p2p_ssd: @@ -1112,6 +1094,45 @@ def __init__(self, cache_config: CacheConfig, model_config: ModelConfig, redis_m # Update initial mempool stats self._update_mempool_metrics() + def _build_cpu_cache_engine(self, + cache_config: CacheConfig, + event_collector: Optional[KVEventCollector]): + """Pick the index engine of the CPU tier. + + A subclass can back the tier with a different engine by overriding this + (see ``flexkv.cache.radix_shmem_planner``). + """ + if cache_config.enable_p2p_cpu: + return HierarchyLRCacheEngine.from_cache_config( + cache_config, self.node_id, DeviceType.CPU, meta=self.redis_meta) + if self.index_accel: + return CacheEngineAccel( + device_type=DeviceType.CPU, + num_total_blocks=cache_config.num_cpu_blocks, + tokens_per_block=cache_config.tokens_per_block, + evict_ratio=self.evict_ratio, + hit_reward_seconds=self.hit_reward_seconds, + evict_start_threshold=self.evict_start_threshold, + eviction_policy=self.eviction_policy, + event_collector=event_collector, + metrics_collector=self._metrics_collector, + protected_threshold=self.protected_threshold, + swa_config=cache_config.swa, + ) + return CacheEngine( + device_type=DeviceType.CPU, + num_total_blocks=cache_config.num_cpu_blocks, + tokens_per_block=cache_config.tokens_per_block, + evict_ratio=self.evict_ratio, + hit_reward_seconds=self.hit_reward_seconds, + evict_start_threshold=self.evict_start_threshold, + eviction_policy=self.eviction_policy, + event_collector=event_collector, + metrics_collector=self._metrics_collector, + protected_threshold=self.protected_threshold, + swa_config=cache_config.swa, + ) + def start(self) -> None: if self.cpu_cache_engine and self.cache_config.enable_p2p_cpu: self.cpu_cache_engine.start() @@ -1153,34 +1174,14 @@ def get(self, namespace: Optional[List[str]] = None, swa_aware: bool = False) \ -> Tuple[TransferOpGraph, np.ndarray, Callable, Dict, int]: - self._check_input(token_ids, token_mask, slot_mapping) - - aligned_length = (token_ids.shape[0] // self.tokens_per_block) * self.tokens_per_block - - aligned_token_ids = token_ids[:aligned_length] - token_mask[aligned_length:] = False - - if aligned_length == 0 or not token_mask.any(): + req = self._prepare_request(token_ids, token_mask, slot_mapping, namespace) + if req.block_end_idx == 0: transfer_graph = TransferOpGraph.create_empty_graph() return_mask = np.zeros_like(token_mask, dtype=np.bool_) callback = partial(self._transfer_callback, node_to_unlock={}, buffer_to_free={}) return transfer_graph, return_mask, callback, {}, -1 - - block_start_idx, block_end_idx = self._get_block_range(token_mask) - # block_end_idx is the block just past the LAST True in token_mask. On the - # plain path the caller marks every non-resident token up to the aligned - # end, so this equals aligned_length // tokens_per_block. On the SWA-aware - # path (swa_aware=True) _get_impl_* clamps the window to usable = min(full, - # swa) after matching, which can end before the aligned length. So the - # invariant is <= (can never exceed the aligned length), not ==. Nothing - # below uses aligned_length; all downstream sizing keys off block_end_idx. - assert block_end_idx <= aligned_length // self.tokens_per_block - gpu_block_ids = self.slot_mapping_to_block_ids(slot_mapping, - self.tokens_per_block)[:block_end_idx-block_start_idx] - - sequence_meta = SequenceMeta(token_ids=aligned_token_ids, - tokens_per_block=self.cache_config.tokens_per_block, - namespace=namespace) + block_start_idx, block_end_idx = req.block_start_idx, req.block_end_idx + gpu_block_ids, sequence_meta = req.gpu_block_ids, req.sequence_meta temp_cache_strategy = resolve_get_cache_strategy( self.use_mooncake_store_backend, temp_cache_strategy) @@ -2048,22 +2049,11 @@ def put(self, temp_cache_strategy: CacheStrategy = DEFAULT_CACHE_STRATEGY, namespace: Optional[List[str]] = None) \ -> Tuple[TransferOpGraph, np.ndarray, Callable, Dict, int]: - self._check_input(token_ids, token_mask, slot_mapping) - # ignore the last incomplete block - aligned_length = (token_ids.shape[0] // self.tokens_per_block) * self.tokens_per_block - aligned_token_ids = token_ids[:aligned_length] - token_mask[aligned_length:] = False - block_start_idx, block_end_idx = self._get_block_range(token_mask) - + req = self._prepare_request(token_ids, token_mask, slot_mapping, namespace) + block_start_idx, block_end_idx = req.block_start_idx, req.block_end_idx # the mask should has a prefix of True assert block_start_idx == 0 - - gpu_block_ids = self.slot_mapping_to_block_ids(slot_mapping, - self.tokens_per_block)[:block_end_idx-block_start_idx] - - sequence_meta = SequenceMeta(token_ids=aligned_token_ids, - tokens_per_block=self.cache_config.tokens_per_block, - namespace=namespace) + gpu_block_ids, sequence_meta = req.gpu_block_ids, req.sequence_meta assert not temp_cache_strategy.ignore_gpu if not self.cache_config.enable_remote or temp_cache_strategy.ignore_remote: @@ -3473,6 +3463,37 @@ def match_all(self, return cpu_matched_result, ssd_matched_result, remote_matched_result + def _prepare_request(self, + token_ids: np.ndarray, + token_mask: np.ndarray, + slot_mapping: np.ndarray, + namespace: Optional[List[str]]) -> RequestWindow: + """Shared GET/PUT prologue. + + Validates the arrays, drops the trailing partial block (``token_mask`` + is cleared past the aligned length in place) and resolves the masked + block window. ``block_end_idx == 0`` means nothing is left to plan. + """ + self._check_input(token_ids, token_mask, slot_mapping) + aligned_length = (token_ids.shape[0] // self.tokens_per_block) * self.tokens_per_block + aligned_token_ids = token_ids[:aligned_length] + token_mask[aligned_length:] = False + block_start_idx, block_end_idx = self._get_block_range(token_mask) + # block_end_idx is the block just past the LAST True in token_mask. On the + # plain path the caller marks every non-resident token up to the aligned + # end, so this equals aligned_length // tokens_per_block. On the SWA-aware + # path (swa_aware=True) _get_impl_* clamps the window to usable = min(full, + # swa) after matching, which can end before the aligned length. So the + # invariant is <= (can never exceed the aligned length), not ==. Nothing + # downstream uses aligned_length; all sizing keys off block_end_idx. + assert block_end_idx <= aligned_length // self.tokens_per_block + gpu_block_ids = self.slot_mapping_to_block_ids( + slot_mapping, self.tokens_per_block)[:block_end_idx - block_start_idx] + sequence_meta = SequenceMeta(token_ids=aligned_token_ids, + tokens_per_block=self.cache_config.tokens_per_block, + namespace=namespace) + return RequestWindow(block_start_idx, block_end_idx, gpu_block_ids, sequence_meta) + def _check_input(self, token_ids: np.ndarray, token_mask: np.ndarray, diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py new file mode 100644 index 000000000..79a83c61d --- /dev/null +++ b/flexkv/cache/radix_shmem_engine.py @@ -0,0 +1,362 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +The radixshmem CPU tier: one process's `shmradix.RadixClient` on the node's +radix-server (index shm + SlotStore + RDMA engine, see +`flexkv.server.shm_radix_bootstrap`). Used only by +`flexkv.cache.radix_shmem_planner.RadixShmemCacheEngine`. + +Contracts: + +1. Publish after transfer: `take()` -> transfer -> `insert()`. A block is + servable by being in the tree, so `insert()` runs from the completion + callback (`StagedRadixInsert`). Slots neither inserted nor recycled leak. +2. A match is a pin: `ShmRadixMatch.release()` must run on every path, or the + prefix stays pinned for the life of the region. +3. `match()` is local only. Peer blocks arrive through `prefetch()` + (`RadixClient.pull_async`), which publishes them into the local tree. + +A region may carry an SWA pool next to FULL. `FULL|SWA` queries return the +joint hit plus the W-block window ending there; SWA slots are addressed by +`component=` and publish only after the Full path. +""" +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Callable, Optional, Sequence + +import numpy as np + +from flexkv.common.debug import flexkv_logger +from flexkv.common.transfer import DeviceType + +if TYPE_CHECKING: # these pull in the C++ extension; keep them off import time + from flexkv.common.block import SequenceMeta + from flexkv.common.config import SWAPoolConfig + from flexkv.integration.dynamo.collector import KVEventCollector + +try: + import shmradix +except ImportError as e: # pragma: no cover + raise ImportError( + "shmradix is not installed; install it from the radixshmem repo " + "(pip install -e radixshmem/python)") from e + +COMPONENT_MASK_FULL = int(shmradix.COMPONENT_MASK_FULL) +COMPONENT_MASK_SWA = int(shmradix.COMPONENT_MASK_SWA) +COMPONENT_FULL = shmradix.ComponentType.FULL +COMPONENT_SWA = shmradix.ComponentType.SWA + + +def _empty_i64() -> np.ndarray: + return np.empty(0, dtype=np.int64) + + +@dataclass +class ShmRadixMatch: + """One local prefix query, pinned until `release()`. + + `local_slots[i]` holds block `i` of the `num_matched_blocks`-block hit. On + a `FULL|SWA` query the hit is the joint one and `swa_slots` cover + `[swa_start, num_matched_blocks)`; both are empty otherwise. + """ + num_matched_blocks: int = 0 + local_slots: np.ndarray = field(default_factory=_empty_i64) + swa_start: int = 0 + swa_slots: np.ndarray = field(default_factory=_empty_i64) + finalize: Optional[Callable[[], None]] = None + + def local_range(self, first: int, last: int) -> np.ndarray: + """Slots of block range [first, last), clipped to the hit.""" + return self.local_slots[first:last] + + def release(self) -> None: + """Drop the pin. Idempotent.""" + finalize, self.finalize = self.finalize, None + if finalize is not None: + finalize() + + +class StagedRadixInsert: + """Staged slots waiting for their transfer. Exactly one of `publish()` + (graph completed) or `abort()` (graph never ran) must run; both then drop + `holds`, the match pin `publish` needs while it inserts. + + A Full+SWA PUT arms two instances, FULL first: SWA's insert refuses paths + the Full tree does not reach yet. + """ + + def __init__(self, + engine: CacheEngineRadixShmem, + sequence_meta: SequenceMeta, + slots: np.ndarray, + path_end: int, + label: str, + holds: Sequence[Callable[[], None]] = (), + component: shmradix.ComponentType = COMPONENT_FULL) -> None: + self._engine = engine + self._sequence_meta = sequence_meta + self._slots = slots + self._path_end = path_end + self._label = label + self._holds = list(holds) + self._component = component + self._settled = False + + def publish(self) -> None: + if self._settled: + return + self._settled = True + try: + self._engine.insert(self._sequence_meta, self._slots, + num_insert_blocks=self._path_end, + component=self._component) + except Exception as e: + flexkv_logger.error( + f"radixshmem {self._label}: insert of {len(self._slots)} " + f"staged slots failed: {e}; returning them to the mempool") + self._engine.recycle(self._slots, component=self._component) + finally: + self._release_holds() + + def abort(self) -> None: + if self._settled: + return + self._settled = True + try: + self._engine.recycle(self._slots, component=self._component) + except Exception as e: + flexkv_logger.error( + f"radixshmem {self._label}: recycle of {len(self._slots)} " + f"staged slots failed: {e}") + finally: + self._release_holds() + + def _release_holds(self) -> None: + holds, self._holds = self._holds, [] + for release in holds: + try: + release() + except Exception as e: # keep releasing: a held ref pins for good + flexkv_logger.error( + f"radixshmem {self._label}: ref release failed: {e}") + + +class CacheEngineRadixShmem: + """One process's attachment to the node's radix-server. Several instances + (one per DP scheduler process) share a server and operate concurrently.""" + + def __init__(self, + shm_name: str, + *, + tokens_per_block: int, + num_total_blocks: int, + peer_enabled: bool = False, + swa_config: Optional[SWAPoolConfig] = None, + event_collector: Optional[KVEventCollector] = None, + metrics_collector=None): + """`shm_name` is the base index name (`radix_index_name`); the server + must be running or starting. `peer_enabled` only takes effect on a + clustered region. `num_total_blocks` is FlexKV's expectation; the + region's capacity is authoritative.""" + from flexkv.server.shm_radix_bootstrap import attach_radix_client + + self.event_collector = event_collector + self._metrics_collector = metrics_collector + cpu_swa = swa_config.for_cache_tier(DeviceType.CPU) if swa_config is not None else None + self.swa_enabled = cpu_swa is not None and cpu_swa.num_slots > 0 + + self._client = attach_radix_client(shm_name) + self._tree = self._client # index ops pass through the client + self.shm_name = self._client.info.index_name # node-suffixed when distributed + self.is_distributed = bool(self._client.is_distributed()) + self.peer_enabled = bool(peer_enabled) and self.is_distributed + if peer_enabled and not self.is_distributed: + flexkv_logger.warning( + f"radixshmem peer reuse is enabled for {self.shm_name} but the " + f"attached region has world_size=1; prefetch stays local-only") + + region_tpb = int(self._client.block_size()) + if region_tpb != int(tokens_per_block): + raise ValueError( + f"radix-server {self.shm_name} has tokens_per_block={region_tpb}, " + f"FlexKV is configured with {tokens_per_block}") + self.tokens_per_block = int(tokens_per_block) + capacity = int(self._client.mempool_total()) + if num_total_blocks > 0 and capacity != int(num_total_blocks): + flexkv_logger.warning( + f"radix-server {self.shm_name} has {capacity} FULL slots, FlexKV " + f"expected {num_total_blocks}; the index is authoritative") + self.num_total_blocks = capacity + + # ---------- attachment / lifecycle ---------- + + @property + def client(self): + return self._client + + @property + def num_free_blocks(self) -> int: + return int(self._tree.mempool_free()) + + def reset(self) -> None: + """Clear the tree. Invalidates outstanding slot ids and matches.""" + self._tree.reset() + + def close(self) -> None: + client, self._client, self._tree = self._client, None, None + if client is not None: + client.close() + + # ---------- queries ---------- + + @staticmethod + def _hashes(sequence_meta: SequenceMeta, query_end: Optional[int]) -> np.ndarray: + sequence_meta.gen_hashes() + hashes = sequence_meta.block_hashes.view(np.uint64) # int64 -> uint64, same width + return hashes if query_end is None else hashes[:query_end] + + def match(self, + sequence_meta: SequenceMeta, + *, + component_mask: int = COMPONENT_MASK_FULL, + query_end: Optional[int] = None) -> ShmRadixMatch: + """Pinned local prefix match. `query_end` caps the queried path so the + SWA window ends where the caller's restore will.""" + hashes = self._hashes(sequence_meta, query_end) + # A refused query has zeroed fields and an unarmed finalize. + qr = self._tree.query(hashes, mask=component_mask, + local_only=True, lock=True) + common_hit = int(qr.common_hit) + + fragments = list(qr.full_fragments) # local_only: at most one + if len(fragments) > 1: + self._finalize_and_raise( + qr, f"radixshmem local query returned {len(fragments)} fragments; " + f"a local_only query yields at most one") + local_slots = (np.asarray(fragments[0][2], dtype=np.int64) + if fragments else _empty_i64()) + if len(local_slots) != common_hit: + self._finalize_and_raise( + qr, f"radixshmem local query covers {len(local_slots)} blocks " + f"of a {common_hit}-block hit") + + swa_slots = np.asarray(qr.swa_slots, dtype=np.int64) + swa_start = int(qr.swa_start) if len(swa_slots) > 0 else 0 + if len(swa_slots) > 0 and swa_start + len(swa_slots) != common_hit: + self._finalize_and_raise( + qr, f"radixshmem returned {len(swa_slots)} SWA slots at " + f"swa_start={swa_start} for a {common_hit}-block joint hit") + + return ShmRadixMatch(num_matched_blocks=common_hit, + local_slots=local_slots, + swa_start=swa_start, + swa_slots=swa_slots, + finalize=qr.finalize) + + @staticmethod + def _finalize_and_raise(qr, message: str) -> None: + if qr.finalize is not None: # never propagate with the prefix pinned + qr.finalize() + raise RuntimeError(message) + + def prefetch(self, + sequence_meta: SequenceMeta, + *, + component_mask: int = COMPONENT_MASK_FULL, + query_end: Optional[int] = None, + timeout_ms: int = 30000) -> Any: + """Start `pull_async` for the prefix; None when the region has no peers. + + The server pulls one peer's run into local slots and the client publishes + it into the local tree on completion. `lock=False`: nothing stays pinned. + `block=False`: a saturated client completes the job with the local hit. + """ + if not self.peer_enabled: + return None + hashes = self._hashes(sequence_meta, query_end) + job = self._client.pull_async(hashes, component_mask, lock=False, + timeout_ms=int(timeout_ms), block=False) + flexkv_logger.debug( + f"radixshmem prefetch on {self.shm_name}: mask={component_mask:#x} " + f"blocks={len(hashes)} local_hit={job.local_hit} " + f"planned_hit={job.planned_hit} job={job.job_id}") + return job + + # ---------- slots ---------- + + def insert(self, + sequence_meta: SequenceMeta, + physical_block_ids: np.ndarray, + num_insert_blocks: int, + component: shmradix.ComponentType = COMPONENT_FULL) -> None: + """Attach transferred slots: `physical_block_ids[i]` is block + `num_insert_blocks - len(physical_block_ids) + i`. Ownership passes to + radixshmem (`auto_recycle=True`); do not recycle these slots again.""" + sequence_meta.gen_hashes() + hashes = sequence_meta.block_hashes.view(np.uint64) + + slots = np.ascontiguousarray(physical_block_ids, dtype=np.int32) + num_slots = len(slots) + if num_slots == 0: + return + + path_end = min(int(num_insert_blocks), len(hashes)) + start = path_end - num_slots + if start < 0: # caller bug; raising keeps slot ownership with the caller + raise ValueError( + f"radixshmem insert of {num_slots} slots overruns the " + f"{path_end}-block path on {self.shm_name}") + + # SWA inserts are right-aligned by radixshmem and refuse a non-zero start. + tree_start = start if component == COMPONENT_FULL else 0 + result = self._tree.insert(hashes[:path_end], slots, start=tree_start, + auto_recycle=True, component=component) + + landed = num_slots - len(result.unused_slots) + if result.error == shmradix.InsertError.FULL_PATH_MISSING: + flexkv_logger.warning( + f"radixshmem {component} insert on {self.shm_name}: full path " + f"[0, {path_end}) was evicted before the window published " + f"(slots were auto-recycled)") + elif result.error != shmradix.InsertError.OK: + flexkv_logger.warning( + f"radixshmem {component} insert on {self.shm_name} returned " + f"{result.error}: {landed}/{num_slots} blocks landed at " + f"start={start} (unused slots were auto-recycled)") + if landed <= 0: + return + + if self.peer_enabled and component == COMPONENT_FULL: + self._tree.flush() # make the new blocks visible cluster-wide + + if (self.event_collector is not None and component == COMPONENT_FULL + and result.error == shmradix.InsertError.OK): + # Error-free, the only unused slots are a redundant prefix, so what + # landed is the tail of the path. + self.event_collector.publish_stored( + block_hashes=sequence_meta.block_hashes[path_end - landed:path_end], + block_size=self.tokens_per_block, + medium="CPU") + + def take(self, + num_required_blocks: int, + component: shmradix.ComponentType = COMPONENT_FULL) -> np.ndarray: + """Allocate up to `num_required_blocks` slots, evicting unpinned LRU + blocks as needed; fewer come back when the pool cannot supply them + (the SWA pool is all-or-none).""" + slots = np.asarray(self._tree.allocate_slots(num_required_blocks, + component=component), + dtype=np.int64) + if (self._metrics_collector is not None and len(slots) > 0 + and component == COMPONENT_FULL): # SWA has its own pool + self._metrics_collector.record_allocation("cpu", len(slots)) + return slots + + def recycle(self, + physical_blocks: np.ndarray, + component: shmradix.ComponentType = COMPONENT_FULL) -> None: + if physical_blocks is None or len(physical_blocks) == 0: + return + self._tree.recycle_slots(np.ascontiguousarray(physical_blocks, dtype=np.int32), + component=component) diff --git a/flexkv/cache/radix_shmem_planner.py b/flexkv/cache/radix_shmem_planner.py new file mode 100644 index 000000000..eaa129d83 --- /dev/null +++ b/flexkv/cache/radix_shmem_planner.py @@ -0,0 +1,580 @@ +# cython: boundscheck=True, wraparound=True +"""GET / PUT / PREFETCH planning on the radixshmem CPU tier. + +`RadixShmemCacheEngine` is `GlobalCacheEngine` with the CPU tier backed by a +radix-server: `CacheEngineRadixShmem`, a `RadixClient` on the shared index and +SlotStore that `shm_radix_bootstrap` brings up. `KVTaskEngine` picks this class +when ``FLEXKV_ENABLE_RADIXSHMEM=1``. + +Why a subclass rather than more branches in `GlobalCacheEngine`: + +* The tree lives in shared memory and admits a block only once it holds data, + so a PUT inserts from the graph-completion callback (`StagedRadixInsert`) + instead of inserting unready nodes at plan time. +* A match is a pinned query (`ShmRadixMatch`), not a locked `RadixNode`; the + pin drops when the graph completes or the plan is aborted. +* Peer reuse is never spliced into a GET. A prefetch (``ignore_gpu``) starts a + `RadixClient.get_async` that pulls a peer's run into the local tree; the plan + has no ops and `KVTaskEngine` completes the task from the job it finds on the + returned `RadixPlanHandle`. + +Not supported here: the SSD and REMOTE tiers, the Redis-backed P2P paths +(``enable_p2p_cpu`` / ``enable_p2p_ssd``) and kv sharing. Peer reuse follows +the radixshmem YAML instead (`RadixShmemConfig.distributed`). +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from functools import partial +from typing import Any, Callable, Dict, List, Optional, Tuple + +import numpy as np +import nvtx + +from flexkv.cache.cache_engine import ( + DEFAULT_CACHE_STRATEGY, + CacheStrategy, + GlobalCacheEngine, + TransferPlanHandle, + _synchronized_cache_tree, +) +from flexkv.cache.radix_shmem_engine import ( + COMPONENT_FULL, + COMPONENT_MASK_FULL, + COMPONENT_MASK_SWA, + COMPONENT_SWA, + CacheEngineRadixShmem, + ShmRadixMatch, + StagedRadixInsert, +) +from flexkv.common.block import SequenceMeta +from flexkv.common.config import CacheConfig, ModelConfig +from flexkv.common.debug import flexkv_logger +from flexkv.common.radixshmem_config import RadixShmemConfig, get_radixshmem_config +from flexkv.common.transfer import ( + DeviceType, + TransferOp, + TransferOpGraph, + TransferType, + add_virtual_op_for_multiple_finished_ops, +) +from flexkv.integration.dynamo.collector import KVEventCollector + +Action = Callable[[], None] + + +@dataclass +class RadixGetPlan: + """What a radixshmem GET or PREFETCH decided. + + ``on_complete`` runs when the graph completes (match-pin release); + ``on_abort`` when the plan is cancelled before its graph ever launched. + A PREFETCH plan has an empty graph and carries the `GetJob` instead: the + pull it started, the local hit it started from and the hit it plans to + reach (blocks). `KVTaskEngine` completes such a task from the job. + """ + + transfer_graph: TransferOpGraph + finished_ops_ids: List[int] = field(default_factory=list) + num_gpu_blocks_to_transfer: int = 0 + on_complete: List[Action] = field(default_factory=list) + on_abort: List[Action] = field(default_factory=list) + prefetch_job: Optional[Any] = None + prefetch_local_hit_blocks: int = 0 + prefetch_planned_hit_blocks: int = 0 + + @classmethod + def empty(cls) -> RadixGetPlan: + return cls(transfer_graph=TransferOpGraph.create_empty_graph()) + + +@dataclass +class RadixPutPlan: + """What a radixshmem PUT decided. ``on_complete`` publishes the staged + slots (`StagedRadixInsert.publish`); ``on_abort`` returns them + (`StagedRadixInsert.abort`). Both drop the match pin.""" + + transfer_graph: TransferOpGraph + finished_ops_ids: List[int] = field(default_factory=list) + num_gpu_blocks_to_transfer: int = 0 + skipped_gpu_blocks: int = 0 + on_complete: List[Action] = field(default_factory=list) + on_abort: List[Action] = field(default_factory=list) + + @classmethod + def empty(cls) -> RadixPutPlan: + return cls(transfer_graph=TransferOpGraph.create_empty_graph()) + + +class RadixPlanHandle(TransferPlanHandle): + """`TransferPlanHandle` plus the prefetch job a PREFETCH plan rides on. + + `KVTaskEngine` reads ``prefetch_job`` off the handle (``getattr``, so the + base handle needs nothing) and completes a job-backed task from the job + rather than from graph completion. + """ + + __slots__ = ("prefetch_job", "prefetch_local_hit_blocks", + "prefetch_planned_hit_blocks") + + def __init__(self, + complete: Action, + abort: Action, + *, + prefetch_job: Optional[Any] = None, + prefetch_local_hit_blocks: int = 0, + prefetch_planned_hit_blocks: int = 0): + super().__init__(complete, abort) + self.prefetch_job = prefetch_job + self.prefetch_local_hit_blocks = prefetch_local_hit_blocks + self.prefetch_planned_hit_blocks = prefetch_planned_hit_blocks + + +def _noop() -> None: + return None + + +def _check_cache_config(cache_config: CacheConfig) -> None: + if not cache_config.enable_cpu: + raise ValueError("radix_shmem needs enable_cpu=True: it backs the CPU tier") + if cache_config.enable_ssd or cache_config.enable_remote: + raise ValueError( + "radix_shmem backs the CPU tier only; enable_ssd and enable_remote " + f"must be off (got enable_ssd={cache_config.enable_ssd}, " + f"enable_remote={cache_config.enable_remote})") + if cache_config.enable_p2p_cpu or cache_config.enable_p2p_ssd or cache_config.enable_kv_sharing: + raise ValueError( + "radix_shmem does its own peer reuse (etcd + RDMA inside the " + "radix-server, on whenever the radixshmem YAML makes the cluster " + "distributed); enable_p2p_cpu / enable_p2p_ssd must be off") + + +class RadixShmemCacheEngine(GlobalCacheEngine): + """`GlobalCacheEngine` whose CPU tier is a radix-server. See the module + docstring for what differs from the built-in planners.""" + + def __init__(self, + cache_config: CacheConfig, + model_config: ModelConfig, + redis_meta=None, + event_collector: Optional[KVEventCollector] = None): + _check_cache_config(cache_config) + self._radix_config: RadixShmemConfig = get_radixshmem_config() + # GetJobs this engine started and has not yet seen finish; pruned on + # every prefetch and used for back-pressure (client.prefetch_max_inflight). + self._prefetch_jobs: List[Any] = [] + super().__init__(cache_config, model_config, redis_meta, event_collector) + + # ------------------------------------------------------------------ tier + + def _build_cpu_cache_engine(self, + cache_config: CacheConfig, + event_collector: Optional[KVEventCollector]): + """Attach to this node's radix-server as a `RadixClient`. + + The server (index + SlotStore + peer transfer) is brought up by the + KVManager bootstrap process or by the operator (`shm_radix_bootstrap`); + the attach waits for it to be ready. + """ + from flexkv.server.shm_radix_bootstrap import radix_index_name + + rcfg = self._radix_config + return CacheEngineRadixShmem( + radix_index_name(rcfg.local_id), + tokens_per_block=cache_config.tokens_per_block, + num_total_blocks=cache_config.num_cpu_blocks, + peer_enabled=rcfg.distributed, + swa_config=cache_config.swa, + event_collector=event_collector, + metrics_collector=self._metrics_collector, + ) + + def _update_mempool_metrics(self) -> None: + if self._metrics_collector is None or self.cpu_cache_engine is None: + return + tier = self.cpu_cache_engine + # A radixshmem tier reports its own counts (no Mempool object). + pool = getattr(tier, "mempool", tier) + self._metrics_collector.update_mempool_stats( + "cpu", pool.num_total_blocks, pool.num_free_blocks) + + # ------------------------------------------------------------- entrances + + @_synchronized_cache_tree + def get(self, + request_id: int, + token_ids: np.ndarray, + token_mask: np.ndarray, + slot_mapping: np.ndarray, + dp_client_id: int, + temp_cache_strategy: CacheStrategy = DEFAULT_CACHE_STRATEGY, + namespace: Optional[List[str]] = None, + swa_aware: bool = False) \ + -> Tuple[TransferOpGraph, np.ndarray, Callable, Dict, int]: + req = self._prepare_request(token_ids, token_mask, slot_mapping, namespace) + if req.block_end_idx == 0: + return self._empty_result(token_mask) + + if temp_cache_strategy.ignore_gpu: + # Prefetch: pull a peer's run into this node's tree. + plan = self._plan_prefetch( + request_id, req.sequence_meta, req.block_end_idx, swa_aware=swa_aware) + else: + plan = self._plan_get( + request_id, req.sequence_meta, req.block_start_idx, req.block_end_idx, + req.gpu_block_ids, dp_client_id, swa_aware=swa_aware) + + transfer_graph, task_end_op_id = add_virtual_op_for_multiple_finished_ops( + plan.transfer_graph, plan.finished_ops_ids, dp_client_id) + + tpb = self.tokens_per_block + return_mask = np.zeros_like(token_mask, dtype=np.bool_) + if plan.prefetch_job is not None: + # The planned pull [local hit, planned hit); KVTaskEngine rewrites + # it to what actually landed when the job completes. + return_mask[plan.prefetch_local_hit_blocks * tpb: + plan.prefetch_planned_hit_blocks * tpb] = True + else: + return_mask[req.block_start_idx * tpb: + (req.block_start_idx + plan.num_gpu_blocks_to_transfer) * tpb] = True + + handle = RadixPlanHandle( + complete=partial(self._run_actions, plan.on_complete, "completion"), + abort=partial(self._run_actions, plan.on_abort, "abort"), + prefetch_job=plan.prefetch_job, + prefetch_local_hit_blocks=plan.prefetch_local_hit_blocks, + prefetch_planned_hit_blocks=plan.prefetch_planned_hit_blocks, + ) + if self._metrics_collector is not None: + self._update_mempool_metrics() + return transfer_graph, return_mask, handle, {}, task_end_op_id + + @_synchronized_cache_tree + def put(self, + request_id: int, + token_ids: np.ndarray, + token_mask: np.ndarray, + slot_mapping: np.ndarray, + dp_client_id: int, + temp_cache_strategy: CacheStrategy = DEFAULT_CACHE_STRATEGY, + namespace: Optional[List[str]] = None) \ + -> Tuple[TransferOpGraph, np.ndarray, Callable, Dict, int]: + req = self._prepare_request(token_ids, token_mask, slot_mapping, namespace) + # the mask should have a prefix of True + assert req.block_start_idx == 0 + assert not temp_cache_strategy.ignore_gpu + if req.block_end_idx == 0: + return self._empty_result(token_mask) + + plan = self._plan_put(request_id, req.sequence_meta, req.block_start_idx, + req.block_end_idx, req.gpu_block_ids, dp_client_id) + + transfer_graph, task_end_op_id = add_virtual_op_for_multiple_finished_ops( + plan.transfer_graph, plan.finished_ops_ids, dp_client_id) + + tpb = self.tokens_per_block + return_mask = np.zeros_like(token_mask, dtype=np.bool_) + mask_lo = (req.block_start_idx + plan.skipped_gpu_blocks) * tpb + return_mask[mask_lo:mask_lo + plan.num_gpu_blocks_to_transfer * tpb] = True + + handle = RadixPlanHandle( + complete=partial(self._run_actions, plan.on_complete, "completion"), + abort=partial(self._run_actions, plan.on_abort, "abort"), + ) + if self._metrics_collector is not None: + self._update_mempool_metrics() + return transfer_graph, return_mask, handle, {}, task_end_op_id + + @staticmethod + def _empty_result(token_mask: np.ndarray) \ + -> Tuple[TransferOpGraph, np.ndarray, Callable, Dict, int]: + return (TransferOpGraph.create_empty_graph(), + np.zeros_like(token_mask, dtype=np.bool_), + RadixPlanHandle(complete=_noop, abort=_noop), + {}, -1) + + @_synchronized_cache_tree + def _run_actions(self, actions: List[Action], what: str) -> None: + """Completion / abort of a plan. Every action runs even if one fails: + a skipped one would leave a pin or staged slots behind for the life of + the region.""" + for action in actions: + try: + action() + except Exception: + flexkv_logger.error( + f"radixshmem plan {what} action failed", exc_info=True) + + # ----------------------------------------------------------------- match + + def _match_cpu(self, + sequence_meta: SequenceMeta, + swa_aware: bool = False, + swa_query_end: Optional[int] = None) -> ShmRadixMatch: + """Local CPU match for the planners (GET and PUT alike). + + radixshmem's `match` never leaves this node: peer blocks reach the local + tree through `_plan_prefetch`, so a GET sees them as an ordinary local + hit once the prefetch has completed. + + ``swa_aware`` turns the match into a joint FULL|SWA query capped at + ``swa_query_end``: its ``num_matched_blocks`` is then the common hit both + components can serve, and the window (``swa_slots``) ends exactly there. + """ + assert self.cpu_cache_engine is not None + if swa_aware: + return self.cpu_cache_engine.match( + sequence_meta, + component_mask=COMPONENT_MASK_FULL | COMPONENT_MASK_SWA, + query_end=swa_query_end, + ) + return self.cpu_cache_engine.match(sequence_meta) + + # ------------------------------------------------------------------- GET + + def _plan_get(self, + request_id: int, + sequence_meta: SequenceMeta, + block_mask_start: int, + block_mask_end: int, + gpu_block_ids: np.ndarray, + dp_client_id: int, + swa_aware: bool = False) -> RadixGetPlan: + """GET: local match, one H2D (plus the SWA H2D chain when SWA-aware). + + Peer blocks are not spliced in. `_plan_prefetch` (sglang's prefetch + hooks run it ahead of scheduling) pulls a peer's run into this node's + tree, so by the time the request is matched here the local hit already + covers them; whatever is not local is a miss. + + Slots join the tree only once they hold data, so nothing here mutates + the tree: the match pin is the only state, released at graph completion + or abort. + """ + nvtx_range = nvtx.start_range( + message=f"CacheEngine.plan_get_radixshmem[{request_id}]", color="cyan") + swa_active = swa_aware and self.swa_op_constructor.enabled + cpu_match = self._match_cpu( + sequence_meta, swa_aware=swa_active, swa_query_end=block_mask_end) + + end = min(cpu_match.num_matched_blocks, block_mask_end) + if end <= block_mask_start: + # Nothing to restore; drop the query's pin now. + cpu_match.release() + if self._metrics_collector is not None and block_mask_end > block_mask_start: + self._metrics_collector.record_cache_miss( + block_mask_end - block_mask_start) + nvtx.end_range(nvtx_range) + return RadixGetPlan.empty() + + if self._metrics_collector is not None: + self._metrics_collector.record_cache_hit("cpu", end - block_mask_start) + if block_mask_end > end: + self._metrics_collector.record_cache_miss(block_mask_end - end) + + transfer_graph = TransferOpGraph() + op_h2d = TransferOp( + graph_id=transfer_graph.graph_id, + transfer_type=TransferType.H2D, + src_block_ids=cpu_match.local_range(block_mask_start, end), + dst_block_ids=gpu_block_ids[:end - block_mask_start], + dp_client_id=dp_client_id, + ) + transfer_graph.add_transfer_op(op_h2d) + finished_ops_ids = [op_h2d.op_id] + + if swa_active and len(cpu_match.swa_slots) > 0: + swa_h2d_id = self.swa_op_constructor.build_get_chain( + transfer_graph, + gpu_slot_ids=np.zeros(len(cpu_match.swa_slots), dtype=np.int64), + cpu_slot_ids=cpu_match.swa_slots, + dp_client_id=dp_client_id, + ) + if swa_h2d_id is not None: + finished_ops_ids.append(swa_h2d_id) + + nvtx.end_range(nvtx_range) + return RadixGetPlan( + transfer_graph=transfer_graph, + finished_ops_ids=finished_ops_ids, + num_gpu_blocks_to_transfer=end - block_mask_start, + on_complete=[cpu_match.release], + on_abort=[cpu_match.release], + ) + + # -------------------------------------------------------------- PREFETCH + + def _prefetch_inflight(self) -> int: + """Peer pulls this engine started that have not finished yet.""" + live = [job for job in self._prefetch_jobs + if not job.done() and not getattr(job, "cancelled", False)] + self._prefetch_jobs = live + return len(live) + + def _plan_prefetch(self, + request_id: int, + sequence_meta: SequenceMeta, + block_mask_end: int, + swa_aware: bool = False) -> RadixGetPlan: + """PREFETCH: start a peer pull. + + `RadixClient.get_async` queries the cluster, stages local slots for a + peer's run, has the radix-server RDMA-read it and publishes the blocks + into this node's tree when the transfer completes. Nothing moves through + the TE, so the plan's graph is empty and the task completes when the job + does (`KVTaskEngine` polls it). A node without peers, or one with too + many pulls in flight, returns an empty plan: the later local GET says + what is here. + """ + assert self.cpu_cache_engine is not None + engine = self.cpu_cache_engine + if not engine.peer_enabled: + return RadixGetPlan.empty() + client_settings = self._radix_config.client + inflight = self._prefetch_inflight() + if inflight >= client_settings.prefetch_max_inflight: + flexkv_logger.debug( + f"radixshmem prefetch {request_id}: {inflight} peer pulls in flight " + f"(limit {client_settings.prefetch_max_inflight}); skipping the peer walk") + return RadixGetPlan.empty() + swa_active = swa_aware and self.swa_op_constructor.enabled + mask = (COMPONENT_MASK_FULL | COMPONENT_MASK_SWA) if swa_active else COMPONENT_MASK_FULL + job = engine.prefetch( + sequence_meta, + component_mask=mask, + query_end=block_mask_end, + timeout_ms=client_settings.prefetch_timeout_ms, + ) + plan = RadixGetPlan.empty() + if job is None: + return plan + self._prefetch_jobs.append(job) + plan.prefetch_job = job + plan.prefetch_local_hit_blocks = int(job.local_hit) + plan.prefetch_planned_hit_blocks = int(job.planned_hit) + if self._metrics_collector is not None and job.planned_hit > job.local_hit: + self._metrics_collector.record_cache_hit("peer", job.planned_hit - job.local_hit) + return plan + + # ------------------------------------------------------------------- PUT + + def _plan_put(self, + request_id: int, + sequence_meta: SequenceMeta, + block_mask_start: int, + block_mask_end: int, + gpu_block_ids: np.ndarray, + dp_client_id: int) -> RadixPutPlan: + """PUT: local only, like every PUT. + + The one difference from ``GlobalCacheEngine._put_impl_local`` is WHEN + the tree learns about the slots. radixshmem accepts them only once they + hold data, so the insert moves into the graph-completion callback; an + abort before launch hands the slots back instead. + + Block index: 0 cpu_tot block_mask_end + GPU : (skipped) | fragment | + | D2H into new slots + CPU : (cached) -+ + """ + assert self.cpu_cache_engine is not None + cpu_engine = self.cpu_cache_engine + cpu_match = self._match_cpu(sequence_meta) + + def _release_match() -> RadixPutPlan: + # Nothing will consume the matched prefix; drop the query's pin now. + cpu_match.release() + return RadixPutPlan.empty() + + num_skipped = len(cpu_match.local_range(block_mask_start, block_mask_end)) + # Window blocks the CPU tier does not already hold. + num_cpu_new = block_mask_end - block_mask_start - num_skipped + # Same policy as _put_impl_local: a fully-matched CPU prefix ends the PUT. + # This also skips the SWA sidecar, so a window lost to SWA-pool eviction + # is not republished until the Full path ages out (accepted limitation). + if num_cpu_new <= 0: + return _release_match() + + cpu_new = cpu_engine.take(num_required_blocks=num_cpu_new) + if len(cpu_new) < num_cpu_new: + flexkv_logger.warning( + f"radixshmem PUT {request_id} skipped: CPU " + f"{len(cpu_new)}/{num_cpu_new} slots available" + ) + cpu_engine.recycle(cpu_new) + if self._metrics_collector is not None: + self._metrics_collector.record_allocation_failure("local") + return _release_match() + + swa_new: Optional[np.ndarray] = None + if self.swa_op_constructor.enabled: + k = min(block_mask_end, self.cache_config.swa.window_blocks) + swa_take = cpu_engine.take(num_required_blocks=k, component=COMPONENT_SWA) + if len(swa_take) == k: + swa_new = swa_take + else: + # All-or-none contract says this is empty; recycle defensively + # in case it ever is not. + cpu_engine.recycle(swa_take, component=COMPONENT_SWA) + flexkv_logger.warning( + f"radixshmem PUT {request_id}: no {k}-slot SWA window " + f"available; storing Full KV only" + ) + + transfer_graph = TransferOpGraph() + finished_ops_ids: List[int] = [] + + fragment_gpu_blocks = gpu_block_ids[num_skipped:] + op_d2h = TransferOp( + graph_id=transfer_graph.graph_id, + transfer_type=TransferType.D2H, + src_block_ids=fragment_gpu_blocks, + dst_block_ids=cpu_new, + dp_client_id=dp_client_id, + ) + transfer_graph.add_transfer_op(op_d2h) + finished_ops_ids.append(op_d2h.op_id) + + if swa_new is not None: + swa_ops = self.swa_op_constructor.build_put_chain( + transfer_graph, + gpu_slot_ids=np.zeros(len(swa_new), dtype=np.int64), + cpu_slot_ids=swa_new, + dp_client_id=dp_client_id, + return_op_ids=True, + ) + assert swa_ops.d2h_id is not None + finished_ops_ids.append(swa_ops.d2h_id) + + on_complete: List[Action] = [] + on_abort: List[Action] = [] + + def _arm(slots: np.ndarray, hold: Optional[Action], label: str, + component=COMPONENT_FULL) -> None: + staged = StagedRadixInsert(engine=cpu_engine, + sequence_meta=sequence_meta, + slots=slots, + path_end=block_mask_end, + label=label, + holds=[] if hold is None else [hold], + component=component) + on_complete.append(staged.publish) + on_abort.append(staged.abort) + + # The match pin travels with the LAST publish: SWA's insert refuses paths + # the Full tree does not reach yet, so it runs after FULL and releases. + _arm(cpu_new, cpu_match.release if swa_new is None else None, + f"PUT {request_id} CPU") + if swa_new is not None: + _arm(swa_new, cpu_match.release, f"PUT {request_id} CPU SWA", + component=COMPONENT_SWA) + + return RadixPutPlan( + transfer_graph=transfer_graph, + finished_ops_ids=finished_ops_ids, + num_gpu_blocks_to_transfer=len(fragment_gpu_blocks), + skipped_gpu_blocks=num_skipped, + on_complete=on_complete, + on_abort=on_abort, + ) diff --git a/flexkv/storage/allocator.py b/flexkv/storage/allocator.py index 7accf7e8d..0f7ea35b6 100644 --- a/flexkv/storage/allocator.py +++ b/flexkv/storage/allocator.py @@ -315,10 +315,59 @@ def _materialize_shareable_hugepage_tensor(path: str, return _wrap_mmap_tensor(mm, aligned, num_elements, dtype, cleanup_path=None) -def materialize_worker_tensor(data: Union[torch.Tensor, HugePageTensorHandle]) -> torch.Tensor: +@dataclass(frozen=True) +class SlotStoreTensorHandle: + """Worker-side handle for one pool of a radixshmem SlotStore. + + The pool is a POSIX shm (or hugetlbfs) mapping owned by the radix-server; + a transfer worker attaches it by name in its own process and views it as a + tensor, the same way ``HugePageTensorHandle`` re-maps a hugetlbfs file. A + tensor built over the mapping cannot be pickled to the worker (that would + copy the bytes), hence the handle. + """ + data_name: str + hugepage_path: str + kind: int # shmradix.ComponentType value (0 = FULL, 1 = SWA) + num_elements: int + dtype: torch.dtype + + def get_tensor(self) -> torch.Tensor: + from shmradix import _data + store = _data.SlotStore.attach(self.data_name, self.hugepage_path, 60000) + return slot_store_pool_tensor(store, self.kind, self.dtype, self.num_elements) + + +def slot_store_pool_tensor(store: Any, kind: Any, dtype: torch.dtype, + num_elements: int) -> torch.Tensor: + """A 1-D tensor over one whole SlotStore pool (``num_slots x stride`` bytes). + + The stride must equal FlexKV's block bytes (the bootstrap picks + ``slot_align`` so), so the view has the layout of a plain CPU buffer. The + store object is pinned on the tensor: dropping the mapping while the + tensor is in use would leave dangling addresses. + """ + view = store.pool_view(kind) + tensor = torch.frombuffer(view, dtype=torch.uint8) + if dtype != torch.uint8: + if tensor.numel() % dtype.itemsize: + raise ValueError( + f"SlotStore pool of {tensor.numel()} bytes is not a multiple of " + f"{dtype} ({dtype.itemsize} bytes)") + tensor = tensor.view(dtype) + if tensor.numel() < num_elements: + raise ValueError( + f"SlotStore pool holds {tensor.numel()} elements of {dtype}, the CPU " + f"layout needs {num_elements}") + tensor = tensor[:num_elements] + tensor._flexkv_slot_store = store # keep the mapping alive with the view + return tensor + + +def materialize_worker_tensor( + data: Union[torch.Tensor, HugePageTensorHandle, SlotStoreTensorHandle]) -> torch.Tensor: if isinstance(data, torch.Tensor): return data - if isinstance(data, HugePageTensorHandle): + if isinstance(data, (HugePageTensorHandle, SlotStoreTensorHandle)): return data.get_tensor() raise TypeError(f"Unsupported worker tensor type: {type(data)}") diff --git a/flexkv/storage/storage_engine.py b/flexkv/storage/storage_engine.py index 96c7a9c3f..4e82a7f35 100644 --- a/flexkv/storage/storage_engine.py +++ b/flexkv/storage/storage_engine.py @@ -9,7 +9,7 @@ from flexkv.common.debug import flexkv_logger from flexkv.common.memory_handle import TensorSharedHandle from flexkv.common.pool import PoolEndpoint, PoolId -from flexkv.common.storage import StorageHandle, KVCacheLayout, KVCacheLayoutType +from flexkv.common.storage import AccessHandleType, StorageHandle, KVCacheLayout, KVCacheLayoutType from flexkv.common.transfer import DeviceType from flexkv.storage.allocator import ( CPUAllocator, @@ -17,6 +17,8 @@ HugePageAllocator, RemoteAllocator, SSDAllocator, + SlotStoreTensorHandle, + slot_store_pool_tensor, ) @@ -78,8 +80,15 @@ def __init__(self, model_config: ModelConfig, cache_config: CacheConfig, num_layers_per_pp_stage: int, - swa_layer_groups: Optional[List[LayerGroupSpec]] = None): - """Initialize storage engine""" + swa_layer_groups: Optional[List[LayerGroupSpec]] = None, + radix_client: Any = None): + """Initialize storage engine. + + ``radix_client`` (a ``shmradix.RadixClient``, radixshmem mode) makes the + CPU FULL / SWA pools views of the radix-server's SlotStore instead of + allocations of this process; see ``_attach_radix_pool``. + """ + self._radix_client = radix_client # One registry, keyed by the endpoint (pool + tier) plus device id. # # It used to be two dicts -- ``_storage_handles`` and @@ -133,11 +142,14 @@ def __init__(self, layer_groups=self._model_config.layer_groups, tp_size=self._model_config.tp_size, ) - self.allocate( - device_type=DeviceType.CPU, - layout=self._cpu_layout, - dtype=buffer_dtype, - ) + if self._radix_client is not None: + self._attach_radix_pool(self._cpu_layout, buffer_dtype, is_swa=False) + else: + self.allocate( + device_type=DeviceType.CPU, + layout=self._cpu_layout, + dtype=buffer_dtype, + ) if self._cache_config.enable_ssd: if not GLOBAL_CONFIG_FROM_ENV.ssd_layout_type == self._cpu_layout.type: @@ -212,15 +224,18 @@ def __init__(self, layer_groups=self._swa_layer_groups, tp_size=self._model_config.tp_size, ) - self.allocate( - device_type=DeviceType.CPU, - layout=self._swa_cpu_layout, - dtype=torch.uint8, - device_id=0, - raw_data=None, - pool_id=PoolId.SWA, - pin_memory=swa_cfg.pin_memory, - ) + if self._radix_client is not None: + self._attach_radix_pool(self._swa_cpu_layout, torch.uint8, is_swa=True) + else: + self.allocate( + device_type=DeviceType.CPU, + layout=self._swa_cpu_layout, + dtype=torch.uint8, + device_id=0, + raw_data=None, + pool_id=PoolId.SWA, + pin_memory=swa_cfg.pin_memory, + ) if self._cache_config.enable_ssd and swa_cfg.num_ssd_slots > 0: @@ -298,6 +313,62 @@ def __init__(self, ) + def _attach_radix_pool(self, + layout: KVCacheLayout, + dtype: torch.dtype, + is_swa: bool) -> None: + """CPU pool from the radix-server's SlotStore (radixshmem mode). + + The FULL pool backs the main KV blocks, the SWA pool the SWA pages; a + slot is one block, so the pool viewed as a tensor has exactly the layout + FlexKV's own allocation would have had. Slot count and stride are + re-checked against ``layout`` here: a mismatch would otherwise become a + misaddressed transfer, not an error. Workers re-attach the pool by name + through the ``SlotStoreTensorHandle`` in ``worker_data``. + """ + from shmradix import ComponentType + from flexkv.server.shm_radix_bootstrap import layout_block_bytes + + kind = ComponentType.SWA if is_swa else ComponentType.FULL + store = self._radix_client.store + if not store.has_pool(kind): + raise ValueError( + f"radix-server SlotStore {store.name} has no {kind.name} pool but the " + f"FlexKV configuration needs one") + pool = store.pool(kind) + block_bytes = layout_block_bytes(layout, dtype) + if int(pool.num_slots) != int(layout.num_block): + raise ValueError( + f"radix-server {kind.name} pool has {pool.num_slots} slots, FlexKV's CPU " + f"layout has {layout.num_block} blocks; the bootstrap and the TE disagree " + f"on the block count") + if int(pool.slot_bytes) != block_bytes: + raise ValueError( + f"radix-server {kind.name} slot stride is {pool.slot_bytes} B, FlexKV's CPU " + f"block is {block_bytes} B; the bootstrap and the TE disagree on the " + f"block layout") + num_elements = layout.get_total_elements() + tensor = slot_store_pool_tensor(store, kind, dtype, num_elements) + handle = StorageHandle( + handle_type=AccessHandleType.TENSOR, + data=tensor, + kv_layout=layout, + dtype=dtype, + worker_data=SlotStoreTensorHandle( + data_name=store.name, + hugepage_path=self._radix_client.info.hugepage_path, + kind=int(kind), + num_elements=num_elements, + dtype=dtype, + ), + ) + pool_id = PoolId.from_is_swa(is_swa) + self._handles[(PoolEndpoint(pool_id, DeviceType.CPU), 0)] = handle + flexkv_logger.info( + f"[StorageEngine] CPU {'SWA' if is_swa else 'KV'} pool = radix-server SlotStore " + f"{store.name} {kind.name} pool: {pool.num_slots} slots x {pool.slot_bytes} B " + f"({pool.num_slots * pool.slot_bytes / 2**30:.2f} GiB)") + def register_gpu_blocks(self, gpu_blocks: List[TensorSharedHandle], gpu_layout: KVCacheLayout, From 227681da40f400f4ae7809d5b4c4d4eeeb2eae24 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:27:40 +0800 Subject: [PATCH 05/21] kvmanager/kvtask: the radixshmem path and node-local DP * KVManager: with enable_radixshmem, partition the op/graph id space per DP client, let local DP client 0 bootstrap the radix regions and the single shm TE subprocess, and have every client attach through the shm channel. * KVTaskEngine: pick RadixShmemCacheEngine, complete job-backed PREFETCH tasks from the RadixClient job instead of from graph completion, and skip the legacy cross-node TransferManagerOnRemote when a node-local shm TE replaces it (ModelConfig.local_dp_size set). Part 5/7 of the radixshmem rebase (see part 1 for provenance). Co-authored-by: Hao Xu Co-authored-by: Iris Ge Co-authored-by: linhu-nv --- flexkv/kvmanager.py | 189 +++++++++++++++++++++++++++++++++++++++++--- flexkv/kvtask.py | 133 +++++++++++++++++++++++++++++-- 2 files changed, 306 insertions(+), 16 deletions(-) diff --git a/flexkv/kvmanager.py b/flexkv/kvmanager.py index fdcebd7f5..baf5496b0 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -26,6 +26,7 @@ from flexkv.server.server import KVServer, DPClient from flexkv.kvtask import KVTaskEngine, KVResponse from flexkv.common.config import ModelConfig, CacheConfig, GLOBAL_CONFIG_FROM_ENV, MooncakeTransferEngineConfig +from flexkv.common.transfer import TransferOpGraph from flexkv.integration.dynamo.collector import KVEventCollector from flexkv.common.debug import eviction_log_aggregator, flexkv_logger from flexkv.cache.redis_meta import RedisMeta @@ -38,7 +39,8 @@ def __init__(self, dp_client_id: int = 0, server_recv_port: str = "", gpu_register_port: str = "", - event_collector: Optional[KVEventCollector] = None): + event_collector: Optional[KVEventCollector] = None, + local_dp_client_id: Optional[int] = None): # Use the curated ``__str__`` summaries. Dataclass repr includes # credential-bearing fields such as ``redis_password``. flexkv_logger.info( @@ -60,15 +62,55 @@ def __init__(self, else: self.gpu_register_port = self.server_recv_port + "_gpu_register" + self.enable_radixshmem = GLOBAL_CONFIG_FROM_ENV.enable_radixshmem + if self.enable_radixshmem and cache_config.enable_remote: + # CacheEngineRadixShmem indexes the CPU tier in shm and reaches + # peers over RDMA; but the 3rd-party (PCFS) tier has its own + # Redis-published index and GET planner. + raise ValueError( + "radix_shmem and enable_remote (3rd-party remote storage) " + "cannot be enabled at the same time" + ) + if self.enable_radixshmem and cache_config.enable_ssd: + raise ValueError( + "radix_shmem backs the CPU tier only (index + SlotStore + peer " + "pull); set ssd_cache_gb=0 / enable_ssd=False" + ) + if self.enable_radixshmem and (cache_config.enable_p2p_cpu + or cache_config.enable_p2p_ssd): + # Peer reuse is the radix-server's (etcd + RDMA), switched on by the + # radixshmem YAML making the cluster distributed; the Redis-backed + # P2P paths these flags select must stay off. + raise ValueError( + "radix_shmem does its own peer reuse; set enable_p2p_cpu=False " + "and enable_p2p_ssd=False (cross-node reuse follows the " + "radixshmem YAML: expected_min_nodes / num_rht_shards)" + ) + # Prefix of this host's radix regions and TE channels: the YAML's + # cluster_id (plus the node name when several nodes share the host). + self._shm_radix_id = None + if self.enable_radixshmem: + from flexkv.common.radixshmem_config import get_radixshmem_config + self._shm_radix_id = get_radixshmem_config().local_id + flexkv_logger.info( f"[KVManager] IPC ports: server_recv_port={self.server_recv_port}, " f"gpu_register_port={self.gpu_register_port}" + ) + if self.enable_radixshmem: + flexkv_logger.info(f"[KVManager] radix_shmem is enabled" + f"[KVManager] shm_radix_id: {self._shm_radix_id}") + # Multi-instance mode also requires server_client_mode - self.server_client_mode = (model_config.dp_size > 1 or - model_config.instance_num > 1 or - GLOBAL_CONFIG_FROM_ENV.server_client_mode) + if self.enable_radixshmem: + # Force server_client_mode False — KVServer is bypassed entirely. + self.server_client_mode = False + else: + self.server_client_mode = (model_config.dp_size > 1 or + model_config.instance_num > 1 or + GLOBAL_CONFIG_FROM_ENV.server_client_mode) self.server_launch_mode = GLOBAL_CONFIG_FROM_ENV.server_launch_mode if self.server_launch_mode not in ("embedded", "external"): raise ValueError( @@ -80,18 +122,40 @@ def __init__(self, "FLEXKV_SERVER_LAUNCH_MODE=external requires server-client mode" ) + self.dp_client_id = dp_client_id + self.local_dp_client_id = ( + dp_client_id if local_dp_client_id is None else local_dp_client_id + ) + flexkv_logger.info( f"[KVManager] instance_num={model_config.instance_num}, dp_size={model_config.dp_size}, " + f"dp_client_id={self.dp_client_id}, " + f"local_dp_client_id={self.local_dp_client_id}, " f"server_client_mode={self.server_client_mode}, " - f"server_launch_mode={self.server_launch_mode}" + f"server_launch_mode={self.server_launch_mode}, " + f"enable_radixshmem={self.enable_radixshmem}" ) self.redis_meta_client = None self.enable_mps = GLOBAL_CONFIG_FROM_ENV.enable_mps self.owns_mps = self.enable_mps and self.server_launch_mode != "external" - - if self.server_client_mode: - if self.server_launch_mode == "embedded" and dp_client_id == 0: + # The embedded radix-server subprocess — only the bootstrap process + # holds this; others have None. + self._shm_radix_server = None + # TE-process handle — only the bootstrap process holds this. + self._shm_te_process = None + # Local KVTaskEngine for the radix-shmem path (per-DP). + self.kv_task_engine = None + self.server_handle = None + + if self.enable_radixshmem: + self._init_radix_shmem_path(event_collector) + elif self.server_client_mode: + # One KVServer per node: with node-local DP the first rank of each + # node owns it, so nodes 1..n-1 get their own server instead of + # waiting on node 0's. Without node-local DP local_dp_client_id + # equals dp_client_id and this is the previous condition. + if self.server_launch_mode == "embedded" and self.local_dp_client_id == 0: self.server_handle = KVServer.create_server(model_config=model_config, cache_config=cache_config, gpu_register_port=self.gpu_register_port, @@ -130,6 +194,103 @@ def __init__(self, event_collector=event_collector, ) + def _init_radix_shmem_path(self, + event_collector: Optional[KVEventCollector]) -> None: + """Initialize the radix-shmem multi-DP path. + + Everything shared by this inference instance's DP processes on the + node — the radix shm regions and the single TE subprocess — is set up + by the node-local bootstrap proc (local DP client 0) only. Every other + proc builds its own KVTaskEngine and + attaches: `CacheEngineRadixShmem` polls for its region, and the TE + channel handle blocks in `ShmControlBlock.wait_ready`. + + Each CE process gets a disjoint graph/op id range so submissions to the + single shared TE never collide. + """ + from flexkv.common.transfer import TransferOp + + TransferOpGraph.set_graph_id_range(self.dp_client_id << 32, + (self.dp_client_id + 1) << 32) + TransferOp.set_op_id_range(self.dp_client_id << 32, + (self.dp_client_id + 1) << 32) + + try: + if self.local_dp_client_id == 0: + self._bootstrap_radix_shmem() + + # KVTaskEngine reads GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and builds a + # RadixShmemCacheEngine (CPU tier = RadixClient on the radix-server). + self.kv_task_engine = KVTaskEngine( + self.model_config, self.cache_config, + self.gpu_register_port, + redis_meta=self.redis_meta_client, + event_collector=event_collector, + shm_te_server_id=self._shm_radix_id, + shm_te_channel_id=self.local_dp_client_id, + ) + except BaseException: + # A failure after the TE / radix-server subprocesses were spawned + # must not leave them running (the TE would wait for GPU + # registrations forever). + self._shutdown_radix_shmem_children() + raise + + def _shutdown_radix_shmem_children(self) -> None: + if self._shm_te_process is not None: + self._shm_te_process.shutdown() + self._shm_te_process = None + if self._shm_radix_server is not None: + self._shm_radix_server.shutdown() + self._shm_radix_server = None + + def _bootstrap_radix_shmem(self) -> None: + """Bootstrap proc (dp 0) only: bring up this node's radix-server (index + + SlotStore + peer transfer) and spawn the shared TE. + + The server is a subprocess (``FLEXKV_RADIX_SERVER_LAUNCH_MODE=embedded``) + or one the operator started (``external``); either way every FlexKV + process attaches by name. Peer reuse needs no Redis address book any + more: the server resolves peers through etcd and pulls their blocks + itself (``RadixClient.pull_async`` from the prefetch path).""" + from flexkv.server.shm_radix_bootstrap import (RadixServerProcess, + build_radix_server_config, + radix_socket_path) + from flexkv.transfer_manager import TransferManagerShmTEProcess + + launch_mode = GLOBAL_CONFIG_FROM_ENV.radix_server_launch_mode + if launch_mode not in ("embedded", "external"): + raise ValueError( + "FLEXKV_RADIX_SERVER_LAUNCH_MODE must be embedded or external, " + f"got {launch_mode!r}" + ) + if launch_mode == "embedded": + server_cfg = build_radix_server_config(self.model_config, self.cache_config) + self._shm_radix_server = RadixServerProcess(server_cfg).start() + self.cache_config.distributed_node_id = int( + self._shm_radix_server.cluster_rank) + flexkv_logger.info( + f"[kv manager] radix-server for {self._shm_radix_id} is up: " + f"cluster rank {self.cache_config.distributed_node_id}" + ) + else: + from flexkv.common.radixshmem_config import get_radixshmem_config + flexkv_logger.info( + f"[kv manager] attaching to an external radix-server at " + f"{get_radixshmem_config().endpoint or radix_socket_path(self._shm_radix_id)}" + ) + + total_clients = self.model_config.total_clients + if self.model_config.local_dp_size is not None: + total_clients = self.model_config.local_dp_size + self._shm_te_process = TransferManagerShmTEProcess( + self.model_config, self.cache_config, + gpu_register_port=self.gpu_register_port, + server_id=self._shm_radix_id, + num_channels=total_clients, + ) + self._shm_te_process.start() + def start(self) -> None: if self.owns_mps: # try to start MPS @@ -161,7 +322,12 @@ def shutdown(self) -> None: self.server_handle.shutdown() self.server_handle = None else: - self.kv_task_engine.shutdown() + if self.kv_task_engine is not None: + self.kv_task_engine.shutdown() + + # Multi-DP radix-shmem teardown — only the bootstrap proc owns these. + # TE first: its workers map the server's SlotStore. + self._shutdown_radix_shmem_children() if self.owns_mps: flexkv_logger.info( @@ -192,6 +358,7 @@ def get_async(self, token_ids=token_ids, slot_mapping=slot_mapping, token_mask=token_mask, + dp_client_id=self.dp_client_id, namespace=namespace, ) return task_id @@ -225,6 +392,7 @@ def get_match(self, token_ids=token_ids, token_mask=token_mask, cpu_only=cpu_only, + dp_client_id=self.dp_client_id, namespace=namespace, swa_aware=swa_aware, ) @@ -250,6 +418,7 @@ def put_async(self, token_ids=token_ids, slot_mapping=slot_mapping, token_mask=token_mask, + dp_client_id=self.dp_client_id, namespace=namespace, ) return task_id @@ -270,6 +439,7 @@ def put_match(self, task_id, mask = self.kv_task_engine.put_match( token_ids=token_ids, token_mask=token_mask, + dp_client_id=self.dp_client_id, namespace=namespace, ) return task_id, mask @@ -299,6 +469,7 @@ def prefetch_async(self, else: task_id = self.kv_task_engine.prefetch_async( token_ids, + dp_client_id=self.dp_client_id, namespace=namespace, swa_aware=swa_aware, ) diff --git a/flexkv/kvtask.py b/flexkv/kvtask.py index 486f2bd2c..bcaa3cfb9 100644 --- a/flexkv/kvtask.py +++ b/flexkv/kvtask.py @@ -1,6 +1,6 @@ import logging import time -from typing import Dict, Optional, List, Union, Tuple +from typing import Any, Dict, Optional, List, Union, Tuple import threading from enum import Enum from dataclasses import dataclass, field, replace @@ -99,6 +99,11 @@ class KVTask: prefetch_has_swa_remote: bool = False prefetch_namespace: Optional[List[str]] = None prefetch_swa_aware: bool = False + # radixshmem prefetch: the RadixClient.pull_async job this task waits on + # (its graph is empty), and the block range it planned to pull. + prefetch_job: Optional[Any] = None + prefetch_local_hit_blocks: int = 0 + prefetch_planned_hit_blocks: int = 0 def is_completed(self) -> bool: return self.status in [TaskStatus.COMPLETED, TaskStatus.CANCELLED, TaskStatus.FAILED] @@ -139,7 +144,9 @@ def __init__(self, cache_config: CacheConfig, gpu_register_port: Optional[str] = None, redis_meta: RedisMeta = None, - event_collector: Optional[KVEventCollector] = None + event_collector: Optional[KVEventCollector] = None, + shm_te_server_id: Optional[str] = None, + shm_te_channel_id: Optional[int] = None, ): if not cache_config.enable_cpu: raise ValueError("enable_cpu must be True") @@ -167,9 +174,35 @@ def __init__(self, f"[KVTaskEngine] topology: {self.model_config}" ) - self.cache_engine = GlobalCacheEngine(cache_config, model_config, redis_meta, event_collector) + # radixshmem prefetch jobs in flight: task_id -> shmradix PullJob. Polled + # in _update_tasks, the thread every other task mutation runs on. + self.prefetch_jobs: Dict[int, Any] = {} + if GLOBAL_CONFIG_FROM_ENV.enable_radixshmem: + # The CPU tier is a radix-server (shared index + SlotStore); its + # planners are a GlobalCacheEngine subclass. + from flexkv.cache.radix_shmem_planner import RadixShmemCacheEngine + self.cache_engine = RadixShmemCacheEngine( + cache_config, model_config, redis_meta, event_collector) + else: + self.cache_engine = GlobalCacheEngine(cache_config, model_config, redis_meta, event_collector) - if not self.model_config.use_trtllm_subprocess: + # Multi-DP shm path: connect this CE to a pre-existing TE process + # via a named ShmChannel rather than spawning a new TE subprocess. + use_shm_te = (shm_te_server_id is not None + and shm_te_channel_id is not None) + if use_shm_te and not self.model_config.use_trtllm_subprocess: + self.transfer_handles = [TransferManagerHandle( + # Left behind by a rename: the sibling "process" branch below + # passes `model_config`, and no *_for_transfer variant exists — + # so the shm-TE path (radix_shmem) NameError'd on first use. + model_config, + self.cache_config, + mode="shm", + gpu_register_port=gpu_register_port, + shm_server_id=shm_te_server_id, + shm_channel_id=shm_te_channel_id, + )] + elif not self.model_config.use_trtllm_subprocess: self.transfer_handles = [TransferManagerHandle( model_config, cache_config, @@ -198,7 +231,9 @@ def __init__(self, ] self.transfer_handles[0]._handle.send_config_to_remotes() - if self.model_config.nnodes > 1: + # A node-local shm TE replaces the legacy cross-node remote manager. + needs_remote_transfer_manager = self.model_config.local_dp_size is None + if self.model_config.nnodes > 1 and needs_remote_transfer_manager: # Bind the handle rather than reading it back with a negative # index: release builds cythonize this module with # wraparound=False, so ``transfer_handles[-1]`` on a list reads off @@ -423,6 +458,13 @@ def create_prefetch_task(self, prefetch_has_swa_remote=prefetch_has_swa_remote, prefetch_namespace=namespace, prefetch_swa_aware=swa_aware) + job = getattr(callback, "prefetch_job", None) + if job is not None: + task = self.tasks[task_id] + task.prefetch_job = job + task.prefetch_local_hit_blocks = int(callback.prefetch_local_hit_blocks) + task.prefetch_planned_hit_blocks = int(callback.prefetch_planned_hit_blocks) + self.prefetch_jobs[task_id] = job self.prefetch_tasks[self._gen_prefetch_key(token_ids, namespace)] = task_id @@ -453,6 +495,7 @@ def _launch_task(self, task_id: int) -> None: transfer_handle.submit(transfer_graph, task_end_op_id=self.tasks[task_id].task_end_op_id) def _update_tasks(self, timeout: float = 0.001) -> None: + self._poll_prefetch_jobs() completed_ops = self._get_completed_ops(timeout) metrics_collector = get_global_collector() for completed_op in completed_ops: @@ -748,6 +791,12 @@ def _cancel_task(self, task_id: int) -> None: if task_id not in self.tasks: return task = self.tasks[task_id] + job = getattr(self, "prefetch_jobs", {}).pop(task_id, None) + if job is not None: + # The pull finishes in the background and still publishes its + # blocks; cancel only drops this task's claim on the result. + job.cancel() + task.prefetch_job = None if not task.is_completed(): # A task whose graph never launched still holds everything its # plan acquired at create time: locked radix nodes and staging @@ -885,8 +934,66 @@ def _process_empty_graph(self, task_id: int) -> None: if task.graph is None: return if task.graph.num_ops == 0: + if task.prefetch_job is not None: + # Nothing for the TE to do, but the job is still moving bytes: + # the task completes from the job, not from the graph. + if task.prefetch_job.done(): + self._complete_prefetch_job(task_id) + return self._mark_completed(task_id) + def _poll_prefetch_jobs(self) -> None: + # getattr: lightweight test/fallback managers built with ``__new__`` + # predate the job table. + jobs = getattr(self, "prefetch_jobs", None) + if not jobs: + return + for task_id in [tid for tid, job in jobs.items() if job.done()]: + self._complete_prefetch_job(task_id) + + def _complete_prefetch_job(self, task_id: int) -> None: + """A radixshmem peer pull finished (or failed, or was refused). The + blocks it fetched are already published in the local tree by the + RadixClient completer; report the pulled range as this prefetch's + return_mask (what sglang books as storage hits) and complete the task.""" + job = self.prefetch_jobs.pop(task_id, None) + task = self.tasks.get(task_id) + if job is None or task is None: + return + result = None + try: + result = job.wait(0) + except Exception as e: # noqa: BLE001 - cancelled, or the completer failed + flexkv_logger.warning( + f"[KVTaskEngine] prefetch task {task_id}: peer pull job " + f"{job.job_id} failed: {e}") + task.transfer_failed = True + tpb = self.cache_config.tokens_per_block + mask = task.return_mask + if isinstance(mask, np.ndarray): + mask[:] = False + if result is not None: + lo = task.prefetch_local_hit_blocks * tpb + hi = min(int(result.common_hit), task.prefetch_planned_hit_blocks) * tpb + if hi > lo: + mask[lo:hi] = True + if result is not None: + result.finalize() # lock=False job: nothing pinned, harmless + planned = task.prefetch_planned_hit_blocks - task.prefetch_local_hit_blocks + flexkv_logger.info( + "[FlexKV-IO] operation=prefetch act=peer_pull status=%s " + "flexkv_task_id=%d job=%d source_rank=%d planned_blocks=%d " + "pulled_blocks=%d bytes=%d local_hit_after=%d", + "success" if result.remote_blocks >= planned else "partial", + task_id, job.job_id, result.source_rank, planned, + result.remote_blocks, result.remote_bytes, result.common_hit) + metrics_collector = get_global_collector() + if metrics_collector is not None and result.remote_blocks > 0: + metrics_collector.record_transfer_completed( + TransferType.PEERH2H.value, int(result.remote_blocks), + int(result.remote_bytes), "get") + self._mark_completed(task_id) + def _get_completed_ops(self, timeout: Optional[float] = None) -> List[CompletedOp]: results = [] # Keep lightweight test/fallback managers created with ``__new__`` @@ -1008,9 +1115,14 @@ def __init__(self, cache_config: CacheConfig, gpu_register_port: Optional[str] = None, redis_meta: Optional[RedisMeta] = None, - event_collector: Optional[KVEventCollector] = None + event_collector: Optional[KVEventCollector] = None, + shm_te_server_id: Optional[str] = None, + shm_te_channel_id: Optional[int] = None, ): - super().__init__(model_config, cache_config, gpu_register_port, redis_meta, event_collector) + super().__init__(model_config, cache_config, gpu_register_port, + redis_meta, event_collector, + shm_te_server_id=shm_te_server_id, + shm_te_channel_id=shm_te_channel_id) self.tracer = FlexKVTracer() self.tracer.trace_config(model_config, cache_config, gpu_layout=None) @@ -1492,6 +1604,13 @@ def reset_cache(self) -> None: # Note: reset_cache() runs on the same thread as the callback dispatch # (_update_tasks), so no lock is needed. We keep the graph_to_task # mapping so a late-completing op still resolves to its task and warns. + prefetch_jobs = getattr(self, "prefetch_jobs", {}) + for task_id, job in list(prefetch_jobs.items()): + job.cancel() + task = self.tasks.get(task_id) + if task is not None: + task.prefetch_job = None + prefetch_jobs.clear() for task_id, task in list(self.tasks.items()): if task.is_completed(): continue # already-fired callbacks are harmless From 4e0717d728acead1f0f8d3eeb1a14c31027adef1 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:27:40 +0800 Subject: [PATCH 06/21] integration: sglang node-local DP and DSv4 SWA sizing; vllm adapter fast paths * sglang connector: node-local DP for every tier, size the host SWA slot from the DSv4 sidecar groups before the TE learns them from the GPU registration, reject a missing dp_rank under radixshmem with dp_size > 1. * vllm adapter: slotted task dataclasses, getattr-based namespace extraction and count_nonzero on the match/put masks. * tests/test_sglang_store_protocol.py: the legacy prefetch-result test now sets _swa_kv_pool on its bare connector like its siblings; prefetch_async reads it since the joint Full+SWA prefetch rule (it failed on the branch too). Part 6/7 of the radixshmem rebase (see part 1 for provenance). Co-authored-by: Hao Xu Co-authored-by: linhu-nv --- flexkv/integration/sglang/connector.py | 192 ++++++++++++--------- flexkv/integration/vllm/vllm_v1_adapter.py | 120 ++++++------- tests/test_sglang_store_protocol.py | 1 + 3 files changed, 169 insertions(+), 144 deletions(-) diff --git a/flexkv/integration/sglang/connector.py b/flexkv/integration/sglang/connector.py index f76508d8a..75031c131 100644 --- a/flexkv/integration/sglang/connector.py +++ b/flexkv/integration/sglang/connector.py @@ -61,6 +61,13 @@ from flexkv.transfer.layer_eventfd import build_layerwise_eventfd_socket_path from flexkv.transfer_manager import TransferManagerOnRemote + +def _radixshmem_distributed() -> bool: + """Whether the radixshmem YAML describes a cluster (peer pulls possible).""" + from flexkv.common.radixshmem_config import get_radixshmem_config + return get_radixshmem_config().distributed + + logger = logging.getLogger(__name__) _SGLANG_REQ_ID_UNSET = object() @@ -167,7 +174,7 @@ def __init__( page_size=self.page_size, tp_rank=tp_rank, pp_rank=pp_rank, - dp_rank=dp_rank if dp_rank is not None else 0, + dp_rank=dp_rank, # None is resolved (or rejected) by the config attn_cp_rank=attn_cp_rank, ) self.model_config = self.flexkv_config.model_config @@ -229,14 +236,21 @@ def __init__( # Heterogeneous groups change the bytes represented by one logical # FlexKV block. Recompute CPU/SSD capacities before KVManager starts. self._apply_layer_groups_for_cache_sizing(kv_caches, indexer_group) + # The radixshmem bootstrap sizes the host SWA slot from the config, before + # the TE learns the DSv4 sidecar groups from the GPU registration. + if self._is_dsv4 and self.cache_config.swa is not None: + _, _, swa_layer_groups, _, _ = self._build_dsv4_swa_registration() + self.cache_config.swa.layer_groups = swa_layer_groups self._label = f"[model_config={self.model_config}, rank_info={self.rank_info}]" # 5. On multi-node setups, every node beyond node 0 needs a # TransferManagerOnRemote process (FlexKV side) before any rank # on that node can register GPU buffers. self._remote_process = None + needs_remote_transfer_manager = self.model_config.local_dp_size is None if ( - self.model_config.nnodes > 1 + needs_remote_transfer_manager + and self.model_config.nnodes > 1 and self.rank_info.node_rank > 0 and self.rank_info.local_rank == 0 ): @@ -257,6 +271,7 @@ def __init__( model_config=self.model_config, cache_config=self.cache_config, dp_client_id=self.rank_info.dp_client_id, + local_dp_client_id=self.rank_info.local_dp_client_id, server_recv_port=self.flexkv_config.server_recv_port, gpu_register_port=self.flexkv_config.gpu_register_port, ) @@ -321,6 +336,9 @@ def __init__( self.cache_config.enable_ssd or self.cache_config.enable_remote or self.cache_config.enable_kv_sharing + # radixshmem cluster: prefetch is where a peer's blocks are pulled + # into this node (RadixClient.pull_async); GET then matches locally. + or (GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and _radixshmem_distributed()) ) self._shutdown_done = False @@ -1399,7 +1417,10 @@ def prefetch_async( if self._sync_ctx.is_sync_leader and self.kv_manager is not None: try: prefetch_result = self.kv_manager.prefetch_async( - token_ids=np.asarray(token_ids, dtype=np.int64) + token_ids=np.asarray(token_ids, dtype=np.int64), + # Same joint Full+SWA rule as lookup_kv: without the SWA + # window a later swa_aware match cannot use the prefix. + swa_aware=self._swa_kv_pool is not None, ) # KVManager currently returns # ``(task_id, actual_prefetch_tokens)`` even though older @@ -2339,82 +2360,11 @@ def _register_standard_to_server( ) logger.info("[FlexKV] Registered KV caches to server %s", self._label) - def _register_dsv4_to_server(self, kv_caches: List[torch.Tensor]) -> None: - """Register DSv4 main KV groups plus SWA/state sidecars.""" - layer_groups: List[LayerGroupSpec] = [] - gpu_layouts: List[KVCacheLayout] = [] - handles_per_group: List[List[torch.Tensor]] = [] - all_buffers: List[torch.Tensor] = [] - - for group in self._dsv4_layer_groups: - buffers = group["buffers"] - sample = buffers[0] - if sample.ndim != 2: - raise RuntimeError( - f"FlexKV DSv4 group {group['name']!r} expects 2D page " - f"buffers, got shape={tuple(sample.shape)}" - ) - if any(buf.shape != sample.shape for buf in buffers): - raise RuntimeError( - f"FlexKV DSv4 group {group['name']!r} has mixed shapes" - ) - sub_page_size = int(group["sub_page_size"]) - if self.page_size % int(group["ratio"]) != 0: - raise RuntimeError( - f"FlexKV page_size={self.page_size} is not divisible by " - f"DSv4 ratio={group['ratio']}" - ) - if sample.shape[1] % sub_page_size != 0: - raise RuntimeError( - f"FlexKV DSv4 group {group['name']!r} page stride " - f"{sample.shape[1]} is not divisible by {sub_page_size}" - ) - head_size = sample.shape[1] // sub_page_size - layer_groups.append( - LayerGroupSpec( - num_layers=len(group["layer_ids"]), - num_kv_heads=1, - head_size=head_size, - layer_indices=list(group["layer_ids"]), - compress_ratio=int(group["ratio"]), - dtype=group["dtype"], - ) - ) - gpu_layouts.append( - KVCacheLayout( - type=KVCacheLayoutType.LAYERFIRST, - num_layer=len(group["layer_ids"]), - num_block=sample.shape[0], - tokens_per_block=sub_page_size, - num_head=1, - head_size=head_size, - kv_dim=self.model_config.kv_dim, - num_kv_heads=self.model_config.num_kv_heads, - ) - ) - handles_per_group.append(list(buffers)) - all_buffers.extend(buffers) - - if len(all_buffers) != len(kv_caches): - raise RuntimeError( - f"FlexKV DSv4 flattened {len(all_buffers)} buffers, expected " - f"{len(kv_caches)}" - ) - - # The primary layout owns the full PP-stage layer-id namespace. Group - # layouts remain local because several groups cover disjoint layer sets. - first_layout = gpu_layouts[0] - primary_layout = KVCacheLayout( - type=first_layout.type, - num_layer=self.rank_info.num_layers_per_pp_stage, - num_block=first_layout.num_block, - tokens_per_block=first_layout.tokens_per_block, - num_head=first_layout.num_head, - head_size=first_layout.head_size, - kv_dim=first_layout.kv_dim, - num_kv_heads=first_layout.num_kv_heads, - ) - + def _build_dsv4_swa_registration(self): + """SWA GPU pool geometry for DSv4: (caches, layout, layer_groups, + gpu_layouts, handles_per_group), all None when the model has no SWA + pool. Shared by the GPU registration and by the radixshmem bootstrap, + which must know the sidecar groups before the TE sees a registration.""" swa_caches = None swa_layout = None swa_layer_groups = None @@ -2494,6 +2444,88 @@ def _register_dsv4_to_server(self, kv_caches: List[torch.Tensor]) -> None: swa_handles_per_group.append(list(state_buffers)) swa_caches.extend(state_buffers) + return (swa_caches, swa_layout, swa_layer_groups, swa_gpu_layouts, + swa_handles_per_group) + + def _register_dsv4_to_server(self, kv_caches: List[torch.Tensor]) -> None: + """Register DSv4 main KV groups plus SWA/state sidecars.""" + layer_groups: List[LayerGroupSpec] = [] + gpu_layouts: List[KVCacheLayout] = [] + handles_per_group: List[List[torch.Tensor]] = [] + all_buffers: List[torch.Tensor] = [] + + for group in self._dsv4_layer_groups: + buffers = group["buffers"] + sample = buffers[0] + if sample.ndim != 2: + raise RuntimeError( + f"FlexKV DSv4 group {group['name']!r} expects 2D page " + f"buffers, got shape={tuple(sample.shape)}" + ) + if any(buf.shape != sample.shape for buf in buffers): + raise RuntimeError( + f"FlexKV DSv4 group {group['name']!r} has mixed shapes" + ) + sub_page_size = int(group["sub_page_size"]) + if self.page_size % int(group["ratio"]) != 0: + raise RuntimeError( + f"FlexKV page_size={self.page_size} is not divisible by " + f"DSv4 ratio={group['ratio']}" + ) + if sample.shape[1] % sub_page_size != 0: + raise RuntimeError( + f"FlexKV DSv4 group {group['name']!r} page stride " + f"{sample.shape[1]} is not divisible by {sub_page_size}" + ) + head_size = sample.shape[1] // sub_page_size + layer_groups.append( + LayerGroupSpec( + num_layers=len(group["layer_ids"]), + num_kv_heads=1, + head_size=head_size, + layer_indices=list(group["layer_ids"]), + compress_ratio=int(group["ratio"]), + dtype=group["dtype"], + ) + ) + gpu_layouts.append( + KVCacheLayout( + type=KVCacheLayoutType.LAYERFIRST, + num_layer=len(group["layer_ids"]), + num_block=sample.shape[0], + tokens_per_block=sub_page_size, + num_head=1, + head_size=head_size, + kv_dim=self.model_config.kv_dim, + num_kv_heads=self.model_config.num_kv_heads, + ) + ) + handles_per_group.append(list(buffers)) + all_buffers.extend(buffers) + + if len(all_buffers) != len(kv_caches): + raise RuntimeError( + f"FlexKV DSv4 flattened {len(all_buffers)} buffers, expected " + f"{len(kv_caches)}" + ) + + # The primary layout owns the full PP-stage layer-id namespace. Group + # layouts remain local because several groups cover disjoint layer sets. + first_layout = gpu_layouts[0] + primary_layout = KVCacheLayout( + type=first_layout.type, + num_layer=self.rank_info.num_layers_per_pp_stage, + num_block=first_layout.num_block, + tokens_per_block=first_layout.tokens_per_block, + num_head=first_layout.num_head, + head_size=first_layout.head_size, + kv_dim=first_layout.kv_dim, + num_kv_heads=first_layout.num_kv_heads, + ) + + (swa_caches, swa_layout, swa_layer_groups, swa_gpu_layouts, + swa_handles_per_group) = self._build_dsv4_swa_registration() + self.tp_client.register_to_server( kv_caches=all_buffers, kv_layout=primary_layout, @@ -2509,7 +2541,7 @@ def _register_dsv4_to_server(self, kv_caches: List[torch.Tensor]) -> None: logger.info( "[FlexKV-DSv4] registered %d main groups, SWA=%s, state_groups=%d", len(layer_groups), - bool(swa_buffers), + swa_caches is not None, len(self._dsv4_state_groups), ) diff --git a/flexkv/integration/vllm/vllm_v1_adapter.py b/flexkv/integration/vllm/vllm_v1_adapter.py index 9d5500e59..f146ae554 100644 --- a/flexkv/integration/vllm/vllm_v1_adapter.py +++ b/flexkv/integration/vllm/vllm_v1_adapter.py @@ -2,7 +2,7 @@ import threading import time from typing import TYPE_CHECKING, Optional, Literal, Iterable, Any, List -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from abc import ABC, abstractmethod import numpy as np @@ -169,7 +169,7 @@ class FlexKVConnectorMetadata(KVConnectorMetadata): default_factory=LayerwiseLoadMetadata) -@dataclass +@dataclass(slots=True) class FlexKVTask(ABC): task_id: int = 0 request: "Request" = 0 @@ -205,7 +205,7 @@ def __str__(self): f"task execute cost {self.task_execute_cost*1000:.2f} ms)") -@dataclass(kw_only=True) +@dataclass(kw_only=True, slots=True) class FlexKVGetTask(FlexKVTask): num_computed_tokens: int num_new_matched_tokens: int @@ -223,7 +223,7 @@ def __str__(self): f"task execute cost {self.task_execute_cost*1000:.2f} ms)") -@dataclass(kw_only=True) +@dataclass(kw_only=True, slots=True) class FlexKVPutTask(FlexKVTask): num_matched_tokens: int num_unmatched_tokens: int @@ -376,49 +376,35 @@ def _extract_namespace(self, request: "Request") -> Optional[List[str]]: """ Extract namespace information from vLLM Request for cache isolation. - This method extracts namespace components from multiple sources in priority order: - 1. lora_request.lora_name: LoRA adapter name for multi-tenant LoRA serving - 2. cache_salt: Explicit cache isolation identifier - 3. namespace_info: User-defined namespace hierarchy (can be list or single value) - - The namespace components are combined to form a hierarchical namespace path, - enabling fine-grained KV cache isolation across different tenants, users, or sessions. - - Args: - request: vLLM Request object containing namespace-related fields - - Returns: - Optional[List[str]]: Ordered list of namespace components forming the hierarchy, - or None if no namespace information is available + Order: lora_request.lora_name → cache_salt → namespace_info (list or single). + Returns None if no namespace information is available. - Example: - If request has lora_name="tenant_A", cache_salt="session_1", - namespace_info=["user_1"], the result will be: - ["tenant_A", "session_1", "user_1"] + Hot path: most workloads have none of the three fields set on the + request, so we read all three with `getattr` (avoids `hasattr` + attribute + lookup) and bail out immediately on the common case. """ - namespace_info = [] + # getattr is faster than hasattr+access pair on the common-empty case. + lora_req = getattr(request, 'lora_request', None) + cache_salt = getattr(request, 'cache_salt', None) + user_namespace = getattr(request, 'namespace_info', None) + # Common case: no isolation fields → single early return. + if lora_req is None and cache_salt is None and user_namespace is None: + return None - if hasattr(request, 'lora_request') and request.lora_request is not None: - lora_id = request.lora_request.lora_name + namespace_info: List[str] = [] + if lora_req is not None: + lora_id = lora_req.lora_name if lora_id is not None: namespace_info.append(str(lora_id)) - - if hasattr(request, 'cache_salt') and request.cache_salt is not None: - cache_salt = request.cache_salt - if cache_salt is not None: - namespace_info.append(str(cache_salt)) - - if hasattr(request, 'namespace_info') and request.namespace_info is not None: - user_namespace = request.namespace_info + if cache_salt is not None: + namespace_info.append(str(cache_salt)) + if user_namespace is not None: if isinstance(user_namespace, list): - namespace_info.extend([str(item) for item in user_namespace]) + namespace_info.extend(str(item) for item in user_namespace) else: namespace_info.append(str(user_namespace)) - if len(namespace_info) == 0: - return None - - return namespace_info + return namespace_info if namespace_info else None def _get_match( self, @@ -446,16 +432,23 @@ def _get_match( if num_tokens_to_get == num_computed_tokens: return -1, 0 - np_token_ids = np.array(token_ids) - np_token_mask = np.ones_like(np_token_ids, dtype=bool) + np_token_ids = np.asarray(token_ids, dtype=np.int64) + # Build mask directly without `np.ones_like + slice False`. With + # `np.empty + 2 slice assigns` we skip the implicit memset-then-zero + # round-trip that `ones_like` does — ~3-5 µs at 4K-token prompts. + np_token_mask = np.empty(num_tokens_to_get, dtype=bool) np_token_mask[:num_computed_tokens] = False + np_token_mask[num_computed_tokens:] = True namespace = self._extract_namespace(request) task_id, matched_mask = self.flexkv_manager.get_match( token_ids=np_token_ids, token_mask=np_token_mask, namespace=namespace, ) - num_new_matched_tokens = matched_mask.sum().item() + # `count_nonzero` on a bool ndarray returns a Python int and is ~2x + # faster than `.sum().item()` because it skips the numpy-scalar + # allocation + `__index__` round-trip. + num_new_matched_tokens = int(np.count_nonzero(matched_mask)) # Auto cancel if not call update_state_after_alloc() match_end_time = time.perf_counter() @@ -646,14 +639,15 @@ def _put_match( if num_tokens_to_put == 0: return -1, 0, 0 - np_token_ids = np.array(token_ids) + np_token_ids = np.asarray(token_ids, dtype=np.int64) namespace = self._extract_namespace(request) task_id, unmatched_mask = self.flexkv_manager.put_match( token_ids=np_token_ids, namespace=namespace, ) - num_unmatched_tokens = unmatched_mask.sum().item() + # See _get_match: count_nonzero on bool is ~2x faster than .sum().item() + num_unmatched_tokens = int(np.count_nonzero(unmatched_mask)) num_matched_tokens = num_tokens_to_put - num_unmatched_tokens # Auto cancel if not need to put. @@ -1309,28 +1303,26 @@ def __init__(self, vllm_config: "VllmConfig", role: "KVConnectorRole", # Track scheduled requests to detect preemptions in build_connector_meta self.previous_scheduled_req_ids: set[str] = set() elif role == KVConnectorRole.WORKER: - # Neither rank is knowable from the config on the worker side: - # * vllm's ParallelConfig has no ``tensor_parallel_rank`` field, so - # the tp_rank read in post_init_from_vllm_config is always 0. - # * the mp executor passes ``local_rank`` to the worker as a kwarg - # and never exports LOCAL_RANK, so the - # ``int(os.environ.get('LOCAL_RANK', -1))`` in - # integration/config.py yields -1 and RankInfo.__post_init__ - # derives local_rank from tp_rank=0 -> 0. - # local_rank is what becomes device_id, so leaving it at 0 makes - # every worker register the same device_id and GPU registration - # never reaches expected_gpus. Recover both from the initialized - # process groups. local_rank must be set explicitly: - # dataclasses.replace re-runs __post_init__, but by then local_rank - # is 0 (not < 0) so it is never re-derived. + # vllm's ParallelConfig has no ``tensor_parallel_rank`` field, so + # the value read in post_init_from_vllm_config is always 0 on every + # worker. Override it here using the initialized TP group rank so + # each worker registers a distinct device_id with FlexKV. + # + # ``local_rank`` has to be reset together with it: RankInfo derives + # local_rank only when it is negative, so a replace() that carried + # the stale 0 over would leave every worker registering + # device_id=0 and the GPU registry would wait forever for the + # TP>1 devices that never arrive. vllm's MultiprocExecutor does + # not export LOCAL_RANK, so the fresh tp_rank is the only source; + # an explicitly exported LOCAL_RANK (torchrun / ray launchers) + # stays authoritative. try: - import dataclasses - from vllm.distributed.parallel_state import (get_tp_group, - get_world_group) - rank_info = dataclasses.replace( - rank_info, - tp_rank=get_tp_group().rank_in_group, - local_rank=get_world_group().local_rank) + overrides: dict[str, int] = { + "tp_rank": get_tp_group().rank_in_group + } + if int(os.environ.get("LOCAL_RANK", -1)) < 0: + overrides["local_rank"] = -1 + rank_info = replace(rank_info, **overrides) except Exception as _e: logger.error( f"FlexKV: could not derive ranks from vllm process groups: " diff --git a/tests/test_sglang_store_protocol.py b/tests/test_sglang_store_protocol.py index 72dffc095..3f72cc9d4 100644 --- a/tests/test_sglang_store_protocol.py +++ b/tests/test_sglang_store_protocol.py @@ -118,6 +118,7 @@ def test_lookup_accepts_sglang_array_token_ids(): def test_prefetch_start_accepts_legacy_manager_result_with_planned_tokens(): connector = FlexKVConnector.__new__(FlexKVConnector) connector._prefetch_enabled = True + connector._swa_kv_pool = None connector.kv_manager = MagicMock() connector.kv_manager.prefetch_async.return_value = (23, 256) connector._sync_ctx = SimpleNamespace( From fdbc40f900104b3e81a5eb63a380afbf30f5e678 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:27:40 +0800 Subject: [PATCH 07/21] tests/docs: radixshmem engine, data-plane and e2e tests; changelog tests/radixshmem/ holds the engine + planner + peer-pull suite and the two e2e scripts (single node, prefetch across nodes); tests/test_shm_channel.py covers the rings. CHANGELOG entry for the radixshmem CPU tier; ignore the mooncake tree cloned by install.sh. Part 7/7 of the radixshmem rebase (see part 1 for provenance). Co-authored-by: Hao Xu Co-authored-by: linhu-nv --- .gitignore | 3 + CHANGELOG.md | 6 + tests/radixshmem/radix_e2e_common.py | 272 +++ .../radixshmem/test_e2e_radix_prefetch_p2p.py | 333 +++ tests/radixshmem/test_e2e_radix_shmem.py | 251 +++ tests/radixshmem/test_radix_shmem_engine.py | 1850 +++++++++++++++++ 6 files changed, 2715 insertions(+) create mode 100644 tests/radixshmem/radix_e2e_common.py create mode 100644 tests/radixshmem/test_e2e_radix_prefetch_p2p.py create mode 100644 tests/radixshmem/test_e2e_radix_shmem.py create mode 100644 tests/radixshmem/test_radix_shmem_engine.py diff --git a/.gitignore b/.gitignore index b0462ccaf..82d03127e 100644 --- a/.gitignore +++ b/.gitignore @@ -91,3 +91,6 @@ ssd_cache*/ benchmarks/nvcomp_benchmarks/inputs benchmarks/nvcomp_benchmarks/runs uv.lock + +# mooncake source tree cloned by install.sh --enable-p2p +.mooncake-build/ diff --git a/CHANGELOG.md b/CHANGELOG.md index 8bc7d486c..0e671ea41 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Feature +Universal: +- radixshmem mode now uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. The radix-server runs as a subprocess of the bootstrap DP (`FLEXKV_RADIX_SERVER_LAUNCH_MODE`). radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). See `docs/radixshmem_integration.md` and `docs/radixshmem_cross_node.md` +- radixshmem mode is configured by one YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `cluster` / `data` / `index` / `server` sections pass through by key to radixshmem's `ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig` (validated against the installed dataclasses, geometry keys rejected), a `client` section holds the prefetch limits; `index.register_chunk_size` defaults to `4096 / tokens_per_block` blocks (one RHT registration chunk per 4096 tokens). The file is global: `cluster.cluster_id` is the only namespace (etcd keys and every shm / socket / TE channel name), node identity derives from `cluster.rpc_interface`. The `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone; `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` remain as per-node overrides for co-located nodes. Examples in `examples/radixshmem_configs/`, reference `docs/radixshmem/config_zh.md`. +- radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles now roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not. Peer reuse in this mode follows the radixshmem YAML (`distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. +- `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`) is trimmed to what `RadixShmemCacheEngine` uses: the `CacheEngineAccel`-compatibility parameters (`device_type`, `evict_ratio`, `evict_start_threshold`, `hit_reward_seconds`, `eviction_policy`, `protected_threshold`, `tokens_per_block=-1`), `take(strict=)`, `match(gpu_matched_blocks=)`, the `mempool` view, `start()`, `store` / `cluster_rank` and the `FLEXKV_TRACE_RADIX_PEER` variable are gone (prefetch logs at debug level; the planner reports mempool metrics itself). + Targeting SGLang: - The native FlexKV backend is available in upstream SGLang `v0.5.16` and later; no patch is required ([sglang#29701](https://github.com/sgl-project/sglang/pull/29701)) - Add DeepSeek-V4 support for heterogeneous C4/C128/indexer KV groups, FullKV + SWA dual caches, attention/indexer compress-state sidecars, and layerwise restore ([#225](https://github.com/taco-project/FlexKV/pull/225)) diff --git a/tests/radixshmem/radix_e2e_common.py b/tests/radixshmem/radix_e2e_common.py new file mode 100644 index 000000000..ca77b9a98 --- /dev/null +++ b/tests/radixshmem/radix_e2e_common.py @@ -0,0 +1,272 @@ +"""Shared pieces of the radixshmem end-to-end tests. + +Used by ``test_e2e_radix_shmem.py`` (one node, one or two DP processes on one +radix-server) and ``test_e2e_radix_prefetch_p2p.py`` (two nodes, one cluster): +GPU block patterns and their byte comparison, request construction, the TP +client subprocess that owns a DP's GPU tensors, KVManager put/get wrappers, and +the RDMA / etcd discovery the cross-node test needs. + +FlexKV is imported lazily inside the functions: the DP / node subprocesses set +their ``FLEXKV_*`` environment before the first ``flexkv`` import, and this +module is re-imported by every spawned child. +""" +from __future__ import annotations + +import contextlib +import glob +import multiprocessing as mp +import os +import shutil +import socket +import subprocess +import tempfile +import time +from typing import List, Optional, Tuple + +import numpy as np +import torch + +TOKENS_PER_BLOCK = 16 +PATTERN_SEED = 0x5EED + + +# ------------------------------------------------------------------ host + +def free_port() -> int: + with contextlib.closing(socket.socket()) as sock: + sock.bind(("127.0.0.1", 0)) + return int(sock.getsockname()[1]) + + +def active_rdma_devices() -> List[str]: + """RDMA devices with at least one ACTIVE port, honoring FLEXKV_TEST_RDMA_DEVICES.""" + override = os.getenv("FLEXKV_TEST_RDMA_DEVICES", "").strip() + names = ([d for d in override.split(",") if d] if override + else sorted(os.path.basename(p) + for p in glob.glob("/sys/class/infiniband/*"))) + active = [] + for name in names: + for state in glob.glob(f"/sys/class/infiniband/{name}/ports/*/state"): + with contextlib.suppress(OSError): + with open(state) as handle: + if "ACTIVE" in handle.read(): + active.append(name) + break + return active + + +def start_private_etcd(prefix: str = "flexkv_etcd_"): + """Start a single-member etcd on free ports. Returns (proc, workdir, + registry) or (None, None, None) when no ``etcd`` binary is on PATH.""" + etcd = shutil.which("etcd") + if not etcd: + return None, None, None + client_port, peer_port = free_port(), free_port() + workdir = tempfile.mkdtemp(prefix=prefix) + proc = subprocess.Popen( + [etcd, "--name", "t", "--data-dir", os.path.join(workdir, "data"), + "--listen-client-urls", f"http://127.0.0.1:{client_port}", + "--advertise-client-urls", f"http://127.0.0.1:{client_port}", + "--listen-peer-urls", f"http://127.0.0.1:{peer_port}", + "--initial-advertise-peer-urls", f"http://127.0.0.1:{peer_port}", + "--initial-cluster", f"t=http://127.0.0.1:{peer_port}"], + stdout=open(os.path.join(workdir, "etcd.log"), "w"), stderr=subprocess.STDOUT) + deadline = time.monotonic() + 15 + while time.monotonic() < deadline: + with contextlib.suppress(OSError): + with socket.create_connection(("127.0.0.1", client_port), timeout=0.5): + return proc, workdir, f"etcd://127.0.0.1:{client_port}" + time.sleep(0.2) + proc.kill() + shutil.rmtree(workdir, ignore_errors=True) + return None, None, None + + +def stop_private_etcd(proc, workdir) -> None: + if proc is not None: + proc.terminate() + with contextlib.suppress(Exception): + proc.wait(10) + if workdir: + shutil.rmtree(workdir, ignore_errors=True) + + +def write_radix_config(workdir: str, config: dict, name: str = "radixshmem.yaml") -> str: + """Write the run's radixshmem YAML (``FLEXKV_RADIXSHMEM_CONFIG_PATH``) and + return its path.""" + import yaml + path = os.path.join(workdir, name) + with open(path, "w") as f: + yaml.safe_dump(config, f) + return path + + +def sweep_radix_files(cluster_id: str) -> None: + """Drop what a run under ``cluster_id`` (the radixshmem namespace; every + shm / socket / IPC name of the run contains it) may have left in shm / tmp.""" + for pattern in (f"/dev/shm/*{cluster_id}*", f"/dev/hugepages/*{cluster_id}*", + f"/tmp/flexkv_{cluster_id}*"): + for stale in glob.glob(pattern): + with contextlib.suppress(OSError): + os.unlink(stale) + + +# ----------------------------------------------------------- GPU bytes + +def block_pattern(layer: int, block_id: int, shape, dtype, + writer: int = 0) -> torch.Tensor: + """Deterministic content for one (layer, block) as written by ``writer``. + + Random rather than structured so a transfer landing on the wrong block or + layer cannot compare equal; generated on the CPU so every process derives + it identically. ``writer`` makes the same block differ between writers, + which is how a byte comparison alone tells whose copy a GET served. + """ + generator = torch.Generator().manual_seed( + PATTERN_SEED + writer * 1_000_003 + block_id * 128 + layer) + return torch.randn(tuple(shape), generator=generator).to(dtype) + + +def write_pattern(gpu_tensors, block_ids, writer: int = 0) -> None: + for layer, tensor in enumerate(gpu_tensors): + for block_id in block_ids: + block = tensor[:, block_id] + block.copy_(block_pattern(layer, int(block_id), block.shape, tensor.dtype, writer)) + torch.cuda.synchronize() + + +def clear_blocks(gpu_tensors, block_ids) -> None: + """Zero the blocks a GET is supposed to fill, so a no-op fails the check.""" + for tensor in gpu_tensors: + for block_id in block_ids: + tensor[:, block_id].zero_() + torch.cuda.synchronize() + + +def mismatched_blocks(gpu_tensors, block_ids, writer: int = 0) -> list: + """(layer, block) pairs whose content is not what ``writer`` wrote.""" + bad = [] + for layer, tensor in enumerate(gpu_tensors): + for block_id in block_ids: + got = tensor[:, block_id].cpu() + want = block_pattern(layer, int(block_id), got.shape, got.dtype, writer) + if not torch.equal(got, want): + bad.append((layer, int(block_id))) + return bad + + +# ------------------------------------------------------------ requests + +def build_request(num_blocks: int, first_block: int, seed: int, + tokens_per_block: int = TOKENS_PER_BLOCK): + """(token_ids, slot_mapping, block_ids) for ``num_blocks`` GPU blocks starting + at ``first_block``; token ids are seeded so two processes agree on them.""" + rng = np.random.default_rng(seed) + block_ids = np.arange(first_block, first_block + num_blocks, dtype=np.int64) + slot_mapping = (np.repeat(block_ids, tokens_per_block) * tokens_per_block + + np.tile(np.arange(tokens_per_block), num_blocks)) + token_ids = rng.integers(0, 32000, size=slot_mapping.shape, dtype=np.int64) + return token_ids, slot_mapping, block_ids + + +def put_prefix(kvm, token_ids, slot_mapping, num_blocks: int, + tokens_per_block: int = TOKENS_PER_BLOCK) -> bool: + """PUT the first ``num_blocks`` blocks; True if the task completed.""" + from flexkv.common.request import KVResponseStatus + num_tokens = num_blocks * tokens_per_block + task_id = kvm.put_async(token_ids=token_ids[:num_tokens], + slot_mapping=slot_mapping[:num_tokens]) + status = kvm.wait([task_id], timeout=120, completely=True) + return all(r.status == KVResponseStatus.SUCCESS for r in status.values()) + + +def get_blocks(kvm, token_ids, slot_mapping, + tokens_per_block: int = TOKENS_PER_BLOCK) -> int: + """One-shot ``get_async`` + wait; matched blocks, 0 if it did not succeed.""" + from flexkv.common.request import KVResponseStatus + task_id = kvm.get_async(token_ids=token_ids, slot_mapping=slot_mapping) + response = kvm.wait([task_id], timeout=120, completely=True)[task_id] + if response.status != KVResponseStatus.SUCCESS or response.return_mask is None: + return 0 + return int(np.count_nonzero(response.return_mask)) // tokens_per_block + + +def match_and_load(kvm, token_ids, slot_mapping, + tokens_per_block: int = TOKENS_PER_BLOCK) -> int: + """The sglang connector's two steps: ``get_match`` then ``launch`` + wait for + the matched prefix. Returns the matched blocks, 0 if the load failed.""" + from flexkv.common.request import KVResponseStatus + task_id, mask = kvm.get_match(token_ids=token_ids) + hit_tokens = int(np.count_nonzero(mask)) + if hit_tokens == 0: + return 0 + kvm.launch([task_id], [slot_mapping[:hit_tokens]]) + response = kvm.wait([task_id], timeout=120, completely=True)[task_id] + if response.status != KVResponseStatus.SUCCESS: + return 0 + return hit_tokens // tokens_per_block + + +def wait_kv_manager_ready(kvm, timeout: float = 180.0) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if kvm.is_ready(): + return + time.sleep(0.1) + raise RuntimeError(f"KVManager not ready in {timeout:.0f}s") + + +# ---------------------------------------------------------- TP client + +def tp_client_proc(server_recv_port: str, dp_client_id: int, device_id: int, + model_config, cache_config, num_gpu_blocks: int, child_conn) -> None: + """Spawn target: owns one DP's GPU tensors, registers them with the TE and + hands their IPC handles back over ``child_conn``; then stays alive so the + tensors do.""" + from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType + from flexkv.common.memory_handle import TensorSharedHandle + from flexkv.server.client import KVTPClient + + tp_client = KVTPClient(server_recv_port, dp_client_id, device_id) + gpu_layout = KVCacheLayout( + type=KVCacheLayoutType.LAYERFIRST, + num_layer=model_config.num_layers, + num_block=num_gpu_blocks, + tokens_per_block=cache_config.tokens_per_block, + num_head=model_config.num_kv_heads // model_config.tp_size, + head_size=model_config.head_size, + kv_dim=model_config.kv_dim, + ) + gpu_blocks = [ + torch.zeros(size=tuple(gpu_layout.kv_shape[1:]), + dtype=model_config.dtype).cuda(device_id) + for _ in range(model_config.num_layers) + ] + tp_client.register_to_server(gpu_blocks, gpu_layout) + child_conn.send([TensorSharedHandle(t) for t in gpu_blocks]) + child_conn.close() + while True: + time.sleep(1) + + +def start_tp_client(kvm, dp_client_id: int, device_id: int, model_config, cache_config, + num_gpu_blocks: int) -> Tuple[mp.Process, list]: + """Start the TP client for ``kvm`` and return (process, this process's view + of its GPU tensors).""" + ctx = mp.get_context("spawn") + parent_conn, child_conn = ctx.Pipe() + proc = ctx.Process( + target=tp_client_proc, + args=(kvm.gpu_register_port, dp_client_id, device_id, model_config, + cache_config, num_gpu_blocks, child_conn), + daemon=True, + ) + proc.start() + gpu_tensors = [handle.get_tensor() for handle in parent_conn.recv()] + return proc, gpu_tensors + + +def stop_tp_client(proc: Optional[mp.Process]) -> None: + if proc is not None: + proc.terminate() + proc.join(timeout=10) diff --git a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py new file mode 100644 index 000000000..bd55a53d3 --- /dev/null +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -0,0 +1,333 @@ +"""End-to-end DATA check for the radixshmem peer path: prefetch pulls a peer's +blocks, GET serves them from the local pool, the GPU holds the peer's bytes. + +Two full FlexKV nodes on one host (two processes, two GPUs), one radixshmem +cluster: + + * each node's KVManager launches its own radix-server (index + SlotStore = + the node's CPU pool + RDMA transfer engine) from one shared YAML + (FLEXKV_RADIXSHMEM_CONFIG_PATH: cluster_id, expected_min_nodes=2, registry, + RDMA devices), told apart by the per-node FLEXKV_RADIX_NODE_NAME override + as co-located nodes are; the two servers rendezvous in one etcd namespace, + get dense cluster ranks and an RHT to route by; + * node 0 PUTs a window of GPU blocks holding a per-block pattern; + * node 1 calls ``KVManager.prefetch_async`` for the same tokens: the index walk + finds the prefix on node 0 over RDMA, node 1's radix-server RDMA-reads the + bytes into node 1's SlotStore and the blocks are published in node 1's tree; + ``try_wait`` reports the pulled range; + * node 1 then does the ordinary local ``get_match`` + ``launch`` + ``wait`` and the + GPU blocks are compared byte for byte with what node 0 wrote. + +A second window checks the extension case: node 1 already holds the first +LOCAL_HEAD_BLOCKS of it (its own bytes), the prefetch pulls only the tail, and +the GET serves head and tail from the right writers. + +Requires >=2 CUDA devices, an ACTIVE RDMA port, a shmradix built with RDMA + +etcd + mooncake, and an etcd (FLEXKV_TEST_RADIX_REGISTRY, or ``etcd`` on PATH +for a private one); skips otherwise. Run inside the container: + + PYTHONPATH=/path/to/FlexKV:/path/to/radixshmem/python \\ + LD_LIBRARY_PATH=$RADIXSHMEM_LIBS:$TORCH_LIB:$LD_LIBRARY_PATH \\ + python -m pytest tests/test_e2e_radix_prefetch_p2p.py -s +""" +from __future__ import annotations + +import multiprocessing as mp +import os +import shutil +import tempfile +import time + +import numpy as np +import pytest +import torch + +from radix_e2e_common import ( + TOKENS_PER_BLOCK, + active_rdma_devices, + build_request, + clear_blocks, + match_and_load, + mismatched_blocks, + put_prefix, + start_private_etcd, + start_tp_client, + stop_private_etcd, + stop_tp_client, + sweep_radix_files, + wait_kv_manager_ready, + write_pattern, + write_radix_config, +) + +WORLD_SIZE = 2 +NUM_GPU_BLOCKS = 128 +NUM_CPU_BLOCKS = 256 +NUM_REQUEST_BLOCKS = 32 +LOCAL_HEAD_BLOCKS = 10 +FIRST_BLOCK = 8 +SECOND_FIRST_BLOCK = 64 +SEED_A = 0x5EED +SEED_B = 0xBEEF + + +def _node_name(rank: int) -> str: + return f"r{rank}" + + +def _prefetch_until(kvm, token_ids, want_pulled_blocks: int, timeout: float = 60.0): + """prefetch_async + try_wait until the pulled range reaches the expectation + (the peer's RHT publication is asynchronous, early rounds may miss). + Returns (pulled_blocks, rounds).""" + from flexkv.common.request import KVResponseStatus + deadline = time.monotonic() + timeout + rounds = 0 + pulled = 0 + while time.monotonic() < deadline: + rounds += 1 + task_id = kvm.prefetch_async(token_ids=token_ids) + response = None + while time.monotonic() < deadline: + done = kvm.try_wait([task_id]) + if task_id in done and done[task_id].status != KVResponseStatus.TIMEOUT: + response = done[task_id] + break + time.sleep(0.02) + if response is None: + break + mask = response.return_mask + pulled = int(np.count_nonzero(mask)) // TOKENS_PER_BLOCK if mask is not None else 0 + if pulled >= want_pulled_blocks: + break + time.sleep(0.5) + return pulled, rounds + + +def _node_proc(rank, gpu_id, cluster_id, config_path, + reader_ready, written, read_done, result_q): + """One FlexKV node: rank 0 writes the windows, rank 1 prefetches and reads.""" + # Before any CUDA context exists: each node drives a different device while + # addressing it as device 0. + os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id) + node_name = _node_name(rank) + # FlexKV's own IPC names are per node too (its radix regions get the + # node name appended through the same override). + recv_port = f"ipc:///tmp/flexkv_{cluster_id}_{node_name}" + os.environ.update({ + "FLEXKV_ENABLE_RADIXSHMEM": "1", + "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, + # The two per-node overrides of the global file: etcd keys membership + # by node identity, which defaults to the bind IP the co-located nodes + # share -- so name each node and give the loopback address explicitly. + "FLEXKV_RADIX_NODE_NAME": node_name, + "FLEXKV_RADIX_RPC_ADDRESS": "127.0.0.1", + "FLEXKV_ENABLE_MPS": "0", + "FLEXKV_SERVER_RECV_PORT": recv_port, + }) + + from flexkv.common.config import CacheConfig, GLOBAL_CONFIG_FROM_ENV, ModelConfig + from flexkv.kvmanager import KVManager + + # Built from env at import time; set the fields that matter explicitly in + # case a parent import happened earlier in this process. + GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True + GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path = config_path + GLOBAL_CONFIG_FROM_ENV.radix_node_name = node_name + GLOBAL_CONFIG_FROM_ENV.radix_rpc_address = "127.0.0.1" + GLOBAL_CONFIG_FROM_ENV.enable_mps = False + GLOBAL_CONFIG_FROM_ENV.server_recv_port = recv_port + + tag = f"[node r{rank}]" + model_config = ModelConfig( + num_layers=2, num_kv_heads=4, head_size=128, + dtype=torch.float16, tp_size=1, dp_size=1, + ) + cache_config = CacheConfig( + tokens_per_block=TOKENS_PER_BLOCK, + enable_cpu=True, enable_ssd=False, enable_remote=False, + num_cpu_blocks=NUM_CPU_BLOCKS, + # Peer reuse follows the radixshmem YAML (expected_min_nodes=2 below), + # not enable_p2p_cpu. + ) + + report = {"rank": rank} + kvm = None + tp_proc = None + try: + kvm = KVManager(model_config, cache_config, dp_client_id=0) + kvm.start() + tp_proc, gpu_tensors = start_tp_client(kvm, 0, 0, model_config, cache_config, + NUM_GPU_BLOCKS) + wait_kv_manager_ready(kvm, timeout=180) + report["node_id"] = int(cache_config.distributed_node_id) + print(f"{tag} READY as cluster rank {report['node_id']}", flush=True) + + tokens_a, slots_a, blocks_a = build_request(NUM_REQUEST_BLOCKS, FIRST_BLOCK, SEED_A) + tokens_b, slots_b, blocks_b = build_request(NUM_REQUEST_BLOCKS, SECOND_FIRST_BLOCK, + SEED_B) + + if rank == 1: + # Window B: this node owns the head before node 0 publishes anything. + write_pattern(gpu_tensors, blocks_b[:LOCAL_HEAD_BLOCKS], writer=1) + report["head_put_ok"] = put_prefix(kvm, tokens_b, slots_b, LOCAL_HEAD_BLOCKS) + print(f"{tag} head put ok={report['head_put_ok']}", flush=True) + reader_ready.set() + if not written.wait(240): + raise TimeoutError("writer did not publish in 240s") + + # 1) Window A, nothing local: the prefetch pulls all of it off node 0, + # then the GET serves it from this node's pool. + pulled, rounds = _prefetch_until(kvm, tokens_a, NUM_REQUEST_BLOCKS) + report["a_pulled_blocks"], report["a_prefetch_rounds"] = pulled, rounds + clear_blocks(gpu_tensors, blocks_a) + report["a_hit_blocks"] = match_and_load(kvm, tokens_a, slots_a) + report["a_mismatched"] = mismatched_blocks(gpu_tensors, blocks_a, writer=0) + print(f"{tag} window A: pulled={pulled} rounds={rounds} " + f"hit={report['a_hit_blocks']} mismatched={len(report['a_mismatched'])}", + flush=True) + + # 2) Window B, local head: the prefetch pulls only the tail. + pulled, rounds = _prefetch_until(kvm, tokens_b, + NUM_REQUEST_BLOCKS - LOCAL_HEAD_BLOCKS) + report["b_pulled_blocks"], report["b_prefetch_rounds"] = pulled, rounds + clear_blocks(gpu_tensors, blocks_b) + report["b_hit_blocks"] = match_and_load(kvm, tokens_b, slots_b) + report["b_head_mismatched"] = mismatched_blocks( + gpu_tensors, blocks_b[:LOCAL_HEAD_BLOCKS], writer=1) + report["b_tail_mismatched"] = mismatched_blocks( + gpu_tensors, blocks_b[LOCAL_HEAD_BLOCKS:], writer=0) + print(f"{tag} window B: pulled={pulled} rounds={rounds} " + f"hit={report['b_hit_blocks']} " + f"head_mismatched={len(report['b_head_mismatched'])} " + f"tail_mismatched={len(report['b_tail_mismatched'])}", flush=True) + read_done.set() + else: + if not reader_ready.wait(240): + raise TimeoutError("reader did not lay down its head in 240s") + write_pattern(gpu_tensors, blocks_a, writer=0) + ok = put_prefix(kvm, tokens_a, slots_a, NUM_REQUEST_BLOCKS) + write_pattern(gpu_tensors, blocks_b, writer=0) + report["put_ok"] = put_prefix(kvm, tokens_b, slots_b, NUM_REQUEST_BLOCKS) and ok + print(f"{tag} put ok={report['put_ok']}", flush=True) + written.set() + # Stay up: the reader's radix-server reads THIS node's SlotStore. + if not read_done.wait(300): + raise TimeoutError("reader did not finish in 300s") + except Exception: + import traceback + report["error"] = traceback.format_exc() + print(f"{tag} FAILED\n{report['error']}", flush=True) + reader_ready.set() + written.set() + read_done.set() + finally: + stop_tp_client(tp_proc) + if kvm is not None: + try: + kvm.shutdown() + except Exception as exc: # noqa: BLE001 + report.setdefault("error", f"shutdown: {exc}") + result_q.put(report) + + +def _run(registry: str, rdma_dev: str) -> dict: + # One etcd namespace (and shm prefix) per run keeps concurrent runs apart. + cluster_id = f"p2p{os.getpid()}" + workdir = tempfile.mkdtemp(prefix="flexkv_radix_p2p_") + config_path = write_radix_config(workdir, { + "cluster": { + "cluster_id": cluster_id, + "expected_min_nodes": WORLD_SIZE, + "registry": registry, + "index_dev": rdma_dev, + "rht_slots_per_bucket": 4, + }, + "data": {"transfer_devices": [rdma_dev], "prefault": False}, + }) + ctx = mp.get_context("spawn") + reader_ready, written, read_done = ctx.Event(), ctx.Event(), ctx.Event() + result_q = ctx.Queue() + procs = [] + reports = {} + try: + for rank in range(WORLD_SIZE): + proc = ctx.Process( + target=_node_proc, + args=(rank, rank, cluster_id, config_path, + reader_ready, written, read_done, result_q), + daemon=False, + ) + proc.start() + procs.append(proc) + deadline = time.monotonic() + 600 + while len(reports) < WORLD_SIZE and time.monotonic() < deadline: + try: + report = result_q.get(timeout=5) + reports[report["rank"]] = report + except Exception: + if not any(proc.is_alive() for proc in procs): + break + finally: + for proc in procs: + proc.join(timeout=30) + if proc.is_alive(): + proc.terminate() + proc.join(timeout=10) + sweep_radix_files(cluster_id) + shutil.rmtree(workdir, ignore_errors=True) + return reports + + +@pytest.fixture +def cluster(): + """(etcd registry, rdma device), skipping when the prerequisites are absent; + starts a private etcd when none is configured.""" + pytest.importorskip("shmradix") + if not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE: + pytest.skip(f"needs {WORLD_SIZE} CUDA devices") + devices = active_rdma_devices() + if not devices: + pytest.skip("no ACTIVE RDMA port (checked /sys/class/infiniband/*)") + registry = os.getenv("FLEXKV_TEST_RADIX_REGISTRY", "") + proc = workdir = None + if not registry: + proc, workdir, registry = start_private_etcd("flexkv_p2p_etcd_") + if not registry: + pytest.skip("set FLEXKV_TEST_RADIX_REGISTRY or put etcd on PATH") + try: + yield registry, devices[0] + finally: + stop_private_etcd(proc, workdir) + + +@pytest.mark.e2e +def test_prefetch_pulls_peer_blocks_over_rdma(cluster): + registry, rdma_dev = cluster + reports = _run(registry, rdma_dev) + + assert len(reports) == WORLD_SIZE, f"only {len(reports)}/{WORLD_SIZE} nodes reported" + errors = {rank: r["error"] for rank, r in reports.items() if "error" in r} + assert not errors, errors + writer, reader = reports[0], reports[1] + assert writer["node_id"] != reader["node_id"], "both nodes got the same cluster rank" + assert writer["put_ok"], "writer's puts did not complete" + assert reader["head_put_ok"], "reader's head put did not complete" + + tail = NUM_REQUEST_BLOCKS - LOCAL_HEAD_BLOCKS + assert reader["a_pulled_blocks"] >= NUM_REQUEST_BLOCKS, \ + f"window A: prefetch pulled {reader['a_pulled_blocks']}/{NUM_REQUEST_BLOCKS}" + assert reader["a_hit_blocks"] == NUM_REQUEST_BLOCKS, \ + f"window A: local GET matched {reader['a_hit_blocks']}/{NUM_REQUEST_BLOCKS}" + assert reader["a_mismatched"] == [], f"window A: wrong bytes in {reader['a_mismatched']}" + assert reader["b_pulled_blocks"] >= tail, \ + f"window B: prefetch pulled {reader['b_pulled_blocks']}/{tail} tail blocks" + assert reader["b_hit_blocks"] == NUM_REQUEST_BLOCKS, \ + f"window B: local GET matched {reader['b_hit_blocks']}/{NUM_REQUEST_BLOCKS}" + assert reader["b_head_mismatched"] == [], \ + f"window B: head does not hold node 1's bytes: {reader['b_head_mismatched']}" + assert reader["b_tail_mismatched"] == [], \ + f"window B: tail does not hold node 0's bytes: {reader['b_tail_mismatched']}" + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-v", "-s"])) diff --git a/tests/radixshmem/test_e2e_radix_shmem.py b/tests/radixshmem/test_e2e_radix_shmem.py new file mode 100644 index 000000000..84883c683 --- /dev/null +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -0,0 +1,251 @@ +"""End-to-end test of FLEXKV_ENABLE_RADIXSHMEM=1 on one node: one or two DP scheduler +processes, one radix-server, one shared transfer engine. + +For every ``dp_size`` in the parametrization: + + * dp0 is the bootstrap process: its KVManager launches the radix-server + (index + SlotStore, the node's CPU KV pool) and spawns the single TE; every + other DP attaches to both by name and feeds the TE over its own shm + channel with a disjoint graph/op id range. The run's namespace is the + ``cluster.cluster_id`` of a small YAML written per run + (FLEXKV_RADIXSHMEM_CONFIG_PATH), which is how a deployment names its + regions too. + * Phase 1: every DP PUTs its own requests concurrently through the shared TE. + * Phase 2 (dp_size > 1): dp0 PUTs a prefix that dp1 then finds with + ``get_match`` -- the shared index is what the radixshmem path exists for. + * Phase 3: every DP writes a rank-specific byte pattern into its GPU blocks, + PUTs them (D2H into the SlotStore), zeroes the GPU blocks, GETs them back + (H2D out of the SlotStore) and compares byte for byte. A transfer that + reached another DP's GPU, or read the wrong SlotStore slot, shows up here + as a mismatch rather than as an error. + +Requires ``dp_size`` CUDA devices; skips otherwise. Run inside the container: + + PYTHONPATH=/path/to/FlexKV:/path/to/radixshmem/python \\ + LD_LIBRARY_PATH=$RADIXSHMEM_LIBS:$TORCH_LIB:$LD_LIBRARY_PATH \\ + python -m pytest tests/test_e2e_radix_shmem.py -s +""" +from __future__ import annotations + +import contextlib +import multiprocessing as mp +import os +import shutil +import tempfile +import time + +import numpy as np +import pytest +import torch + +from radix_e2e_common import ( + TOKENS_PER_BLOCK, + build_request, + clear_blocks, + get_blocks, + mismatched_blocks, + start_tp_client, + stop_tp_client, + sweep_radix_files, + wait_kv_manager_ready, + write_pattern, + write_radix_config, +) + +NUM_GPU_BLOCKS = 256 +NUM_CPU_BLOCKS = 4096 +BLOCK_PER_REQUEST = 32 +# A fixed prefix that dp0 writes and dp1 later looks up across the shared index. +SHARED_SEED = 0x5EED +SHARED_START_BLOCK = 128 +# GPU blocks phase 3 owns, past the ranges phases 1 and 2 use. +ROUNDTRIP_START_BLOCK = 192 + + +def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, + barrier, result_q) -> None: + """Full lifecycle of one DP scheduler process.""" + # Before the first flexkv import: GLOBAL_CONFIG_FROM_ENV is read at import. + # All DP procs share one TE, so they must agree on server_recv_port (and + # therefore on the gpu_register_port the TE listens on). + recv_port = f"ipc:///tmp/flexkv_{server_id}" + os.environ.update({ + "FLEXKV_ENABLE_RADIXSHMEM": "1", + "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, + "FLEXKV_ENABLE_MPS": "0", + "FLEXKV_SERVER_RECV_PORT": recv_port, + }) + + from flexkv.common.config import CacheConfig, GLOBAL_CONFIG_FROM_ENV, ModelConfig + from flexkv.common.request import KVResponseStatus + from flexkv.kvmanager import KVManager + + GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True + GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path = config_path + GLOBAL_CONFIG_FROM_ENV.enable_mps = False + GLOBAL_CONFIG_FROM_ENV.server_recv_port = recv_port + + model_config = ModelConfig( + num_layers=2, num_kv_heads=4, head_size=128, + dtype=torch.float16, tp_size=1, dp_size=dp_size, + ) + cache_config = CacheConfig( + tokens_per_block=TOKENS_PER_BLOCK, enable_cpu=True, enable_ssd=False, + num_cpu_blocks=NUM_CPU_BLOCKS, + ) + + tag = f"[dp{dp_client_id}]" + report = {"dp": dp_client_id} + kvm = None + tp_proc = None + try: + kvm = KVManager(model_config, cache_config, dp_client_id=dp_client_id) + kvm.start() + # Each DP drives its own GPU (device id = dp id; the TE opens every + # DP's IPC handles because total_gpus > 1 clears CUDA_VISIBLE_DEVICES + # for it when dp_size > 1). + tp_proc, gpu_tensors = start_tp_client( + kvm, dp_client_id, dp_client_id, model_config, cache_config, NUM_GPU_BLOCKS) + wait_kv_manager_ready(kvm, timeout=120) + print(f"{tag} READY", flush=True) + + # --- Phase 1: private PUTs, concurrently on every DP. --- + own = [] + for i in range(3): + tok, slot, _blk = build_request( + BLOCK_PER_REQUEST, i * BLOCK_PER_REQUEST, seed=1000 * dp_client_id + i) + own.append(kvm.put_async(token_ids=tok, slot_mapping=slot)) + statuses = kvm.wait(own, timeout=60, completely=True) + report["private_puts"] = ( + sum(1 for s in statuses.values() if s.status == KVResponseStatus.SUCCESS), + len(own)) + print(f"{tag} private puts: {report['private_puts']}", flush=True) + + # --- Phase 2: dp0 PUTs a shared prefix; dp1 finds it in the shared index. --- + report["shared_hit_blocks"] = -1 + if dp_size > 1: + tok, slot, _blk = build_request(BLOCK_PER_REQUEST, SHARED_START_BLOCK, + seed=SHARED_SEED) + if dp_client_id == 0: + task_id = kvm.put_async(token_ids=tok, slot_mapping=slot) + statuses = kvm.wait([task_id], timeout=60, completely=True) + report["shared_put_ok"] = all( + s.status == KVResponseStatus.SUCCESS for s in statuses.values()) + barrier.wait(120) # release dp1 to look it up + else: + barrier.wait(120) # until dp0 finished the shared put + # Index visibility is asynchronous after the store: poll. + hit = 0 + for _ in range(50): + _tid, mask = kvm.get_match(token_ids=tok) + hit = (int(np.count_nonzero(mask)) // TOKENS_PER_BLOCK + if mask is not None else 0) + if hit > 0: + break + time.sleep(0.2) + report["shared_hit_blocks"] = hit + print(f"{tag} cross-DP match hit_blocks={hit}", flush=True) + + # --- Phase 3: byte round trip through this DP's own GPU. --- + tok, slot, blk = build_request(BLOCK_PER_REQUEST, ROUNDTRIP_START_BLOCK, + seed=7000 + dp_client_id) + # The DPs use the same block ids on different devices, so the phases are + # kept in lockstep: a stray transfer then lands on a block under check. + write_pattern(gpu_tensors, blk, writer=dp_client_id) + barrier.wait(120) + task_id = kvm.put_async(token_ids=tok, slot_mapping=slot) + report["roundtrip_put_ok"] = all( + s.status == KVResponseStatus.SUCCESS + for s in kvm.wait([task_id], timeout=60, completely=True).values()) + barrier.wait(120) + clear_blocks(gpu_tensors, blk) + barrier.wait(120) + report["roundtrip_hit_blocks"] = get_blocks(kvm, tok, slot) + barrier.wait(120) + report["roundtrip_mismatched"] = mismatched_blocks(gpu_tensors, blk, writer=dp_client_id) + print(f"{tag} roundtrip: put_ok={report['roundtrip_put_ok']} " + f"hit_blocks={report['roundtrip_hit_blocks']} " + f"mismatched={len(report['roundtrip_mismatched'])}", flush=True) + except Exception: + import traceback + report["error"] = traceback.format_exc() + print(f"{tag} FAILED\n{report['error']}", flush=True) + with contextlib.suppress(Exception): + barrier.abort() # unblock the other DPs' waits + finally: + stop_tp_client(tp_proc) + if kvm is not None: + try: + kvm.shutdown() + except Exception as exc: # noqa: BLE001 + report.setdefault("error", f"shutdown: {exc}") + result_q.put(report) + + +def _run(dp_size: int) -> dict: + server_id = f"e2e{dp_size}dp_{os.getpid()}" + workdir = tempfile.mkdtemp(prefix="flexkv_radix_e2e_") + config_path = write_radix_config(workdir, {"cluster": {"cluster_id": server_id}, + "data": {"prefault": False}}) + ctx = mp.get_context("spawn") + barrier = ctx.Barrier(dp_size) + result_q = ctx.Queue() + procs = [ + ctx.Process(target=_dp_proc, + args=(dp, dp_size, server_id, config_path, barrier, result_q), + daemon=False) + for dp in range(dp_size) + ] + reports = {} + try: + for proc in procs: + proc.start() + deadline = time.monotonic() + 420 + while len(reports) < dp_size and time.monotonic() < deadline: + try: + report = result_q.get(timeout=5) + reports[report["dp"]] = report + except Exception: + if not any(proc.is_alive() for proc in procs): + break + finally: + for proc in procs: + proc.join(timeout=30) + if proc.is_alive(): + proc.terminate() + proc.join(timeout=10) + sweep_radix_files(server_id) + shutil.rmtree(workdir, ignore_errors=True) + return reports + + +@pytest.mark.e2e +@pytest.mark.parametrize("dp_size", [1, 2]) +def test_radix_shmem_put_get_roundtrip(dp_size): + pytest.importorskip("shmradix") + if not torch.cuda.is_available() or torch.cuda.device_count() < dp_size: + pytest.skip(f"needs {dp_size} CUDA device(s)") + + reports = _run(dp_size) + + assert len(reports) == dp_size, f"only {len(reports)}/{dp_size} DP processes reported" + errors = {dp: r["error"] for dp, r in reports.items() if "error" in r} + assert not errors, errors + for dp, report in sorted(reports.items()): + ok, total = report["private_puts"] + assert ok == total, f"dp{dp}: private puts {ok}/{total}" + assert report["roundtrip_put_ok"], f"dp{dp}: round-trip put did not complete" + assert report["roundtrip_hit_blocks"] == BLOCK_PER_REQUEST, \ + f"dp{dp}: round-trip GET matched {report['roundtrip_hit_blocks']}/{BLOCK_PER_REQUEST}" + assert report["roundtrip_mismatched"] == [], \ + f"dp{dp}: {len(report['roundtrip_mismatched'])} block(s) hold KV that is not its " \ + f"own -- its transfers reached another DP's GPU or the wrong SlotStore slot" + if dp_size > 1: + assert reports[0].get("shared_put_ok"), "dp0's shared-prefix put did not complete" + assert reports[1]["shared_hit_blocks"] == BLOCK_PER_REQUEST, \ + f"cross-DP prefix sharing: dp1 matched {reports[1]['shared_hit_blocks']}/" \ + f"{BLOCK_PER_REQUEST} blocks of dp0's prefix" + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-v", "-s"])) diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py new file mode 100644 index 000000000..da3336103 --- /dev/null +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -0,0 +1,1850 @@ +"""Tests for the radixshmem CPU tier: the engine on a radix-server, the data +plane FlexKV maps as its CPU pool, the planners, and the peer pull. +Skipped if `shmradix` is missing. + +Four parts, one file: + + Part 1 — `CacheEngineRadixShmem` semantics against an in-process + `shmradix.RadixServer` (index + SlotStore): take / insert / match / recycle, + insert-after-transfer publication, lock vs. eviction, SWA windows, and the + fact that a standalone region has no peer to prefetch from. + Part 1b — the data plane: the SlotStore pool viewed as FlexKV's CPU tensor, + the exact-stride geometry the bootstrap derives from the configuration, and + the embedded radix-server subprocess. + Part 2 — `GlobalCacheEngine.get()/put()` planning on the radixshmem backend, + driven by synthetic matches (no region): the local GET is one H2D, the + prefetch plan carries a `pull_async` job, the PUT arms the deferred insert; + plus `KVTaskEngine` completing a job-backed prefetch task. + Part 3 — a real 2-node radixshmem cluster over RDMA in two spawned + processes: node 0 publishes a prefix with bytes, node 1 prefetches it (the + server pulls the bytes) and then matches it locally, byte for byte. + +Import-time notes: + * Parts 1 and 3 use a duck-typed fake `SequenceMeta` and side-load + `radix_shmem_engine.py`, so they do not pull in `flexkv.c_ext` (CUDA). + * Part 2 needs the real `GlobalCacheEngine` (and therefore `c_ext`), so it + imports it lazily and skips instead of breaking collection for Parts 1/3. + * Part 3 needs an ACTIVE RDMA device, a shmradix built WITH RDMA + etcd + + mooncake, and an etcd (FLEXKV_TEST_RADIX_REGISTRY, or an `etcd` binary on + PATH to start a private one); it is gated behind FLEXKV_RUN_RADIX_PEER_TEST=1. +""" +from __future__ import annotations + +import contextlib +import copy +import glob +import importlib.util +import multiprocessing as mp +import os +import shutil +import socket +import subprocess +import sys +import tempfile +import time +import traceback +from dataclasses import dataclass +from types import SimpleNamespace + +import numpy as np +import pytest + +try: + import shmradix +except ImportError as exc: + # Not importorskip: a shmradix built before the current API (stale `_core.so` + # next to a newer `__init__.py`) raises ImportError rather than + # ModuleNotFoundError, and pytest >= 8.2 only skips on the latter — which + # would abort collection for the whole suite instead of skipping this file. + pytest.skip(f"shmradix unusable ({exc}); rebuild the extension", + allow_module_level=True) + +for _name in ("RadixServer", "RadixServerConfig", "IndexConfig", "DataPlaneConfig"): + if not hasattr(shmradix, _name): + pytest.skip(f"shmradix lacks {_name}: needs the RadixServer/RadixClient surface", + allow_module_level=True) + + +def _load_module_direct(name: str, path: str): + """Load a module by file path, bypassing parent package __init__. + + `flexkv/cache/__init__.py` imports `flexkv.c_ext`, which links libcudart. + Side-load `radix_shmem_engine` directly so the test runs on CPU-only hosts. + """ + spec = importlib.util.spec_from_file_location(name, path) + module = importlib.util.module_from_spec(spec) + # `@dataclass` looks up the module in sys.modules during class + # construction; register before exec_module. + sys.modules[name] = module + try: + spec.loader.exec_module(module) + except Exception: + sys.modules.pop(name, None) + raise + return module + + +_FLEXKV_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, os.pardir)) +_engine_mod = _load_module_direct( + "_radix_shmem_engine_test", + os.path.join(_FLEXKV_ROOT, "flexkv", "cache", "radix_shmem_engine.py"), +) +CacheEngineRadixShmem = _engine_mod.CacheEngineRadixShmem +# The planner only duck-types the match, so the side-loaded class is as good as +# the one `flexkv.cache.cache_engine` imports — and it needs no c_ext. +ShmRadixMatch = _engine_mod.ShmRadixMatch + +# Pure-Python (no c_ext): the bootstrap (server config, geometry, attach) and +# the transfer enums. +from flexkv.common.config import GLOBAL_CONFIG_FROM_ENV # noqa: E402 +from flexkv.common.radixshmem_config import ( # noqa: E402 + RadixShmemConfigError, load_radixshmem_config, set_radixshmem_config) +from flexkv.common.transfer import TransferType # noqa: E402 +from flexkv.server import shm_radix_bootstrap as bootstrap # noqa: E402 + +FULL = shmradix.ComponentType.FULL +_SWA = shmradix.ComponentType.SWA +SLOT_BYTES = 256 # bytes of one test block in the SlotStore + + +@dataclass +class FakeSeq: + """Duck-typed SequenceMeta for the engine API: only `block_hashes` and + `gen_hashes()` are read by `CacheEngineRadixShmem`.""" + block_hashes: np.ndarray + tokens_per_block: int = 4 + + @property + def num_blocks(self) -> int: + return len(self.block_hashes) + + def gen_hashes(self) -> None: + # already populated + pass + + +def _hashes(seed: int, num: int) -> np.ndarray: + """Deterministic, distinct int64 hashes.""" + rng = np.random.default_rng(seed) + return rng.integers(low=1, high=2**62, size=num, dtype=np.int64) + + +def _sweep_region(name: str, data_name: str | None = None) -> None: + """Drop the shm objects and the socket a previous run may have left.""" + base = name.lstrip("/").replace("/", "_") + paths = [f"/dev/shm/{base}", f"/dev/shm/{base}.sock", + f"/dev/shm/{(data_name or name + '_data').lstrip('/')}"] + for root in ("/dev/hugepages",): + paths.append(f"{root}/{base}") + for path in paths: + with contextlib.suppress(FileNotFoundError, IsADirectoryError): + os.remove(path) + + +def _server_config(name: str, blocks: int, tokens_per_block: int, + swa_slots: int = 0, window_blocks: int = 0, + slot_bytes: int = SLOT_BYTES): + """A standalone data-mode server: one slot per block, stride == slot bytes.""" + align = bootstrap.slot_align_for(slot_bytes) + return shmradix.RadixServerConfig( + index=shmradix.IndexConfig( + name=name, tokens_per_block=tokens_per_block, full_slots=blocks, + swa_slots=swa_slots, swa_window_blocks=window_blocks), + data=shmradix.DataPlaneConfig( + data_bytes=(blocks + swa_slots) * slot_bytes, full_slot_bytes=slot_bytes, + swa_slot_bytes=slot_bytes if swa_slots else 0, slot_align=align, + prefault=False), + ) + + +class _Env: + """Owns the in-process servers and engines a test creates; closes them in + reverse order at teardown (engine first, it maps the server's regions).""" + + def __init__(self) -> None: + self._stack = [] + + def server(self, cfg): + _sweep_region(cfg.index.name, cfg.resolved_data_name) + server = shmradix.RadixServer(cfg).start() + self._stack.append(server.close) + return server + + def engine(self, name: str, **kwargs) -> CacheEngineRadixShmem: + engine = CacheEngineRadixShmem(name, **kwargs) + self._stack.append(engine.close) + return engine + + def make(self, name: str, blocks: int = 2000, tokens_per_block: int = 4, **engine_kwargs): + cfg = _server_config(name, blocks, tokens_per_block) + server = self.server(cfg) + engine = self.engine(name, num_total_blocks=blocks, + tokens_per_block=tokens_per_block, **engine_kwargs) + return engine, server + + def close(self) -> None: + while self._stack: + with contextlib.suppress(Exception): + self._stack.pop()() + + +def _radix_config(**cluster): + """The all-defaults radixshmem configuration (standalone, socket derived + from the index name) with ``cluster`` keys changed.""" + return load_radixshmem_config(None).replace_cluster(**cluster) + + +@pytest.fixture +def env(): + set_radixshmem_config(_radix_config()) + e = _Env() + try: + yield e + finally: + e.close() + set_radixshmem_config(None) + + +# ============================================================================= +# Part 1 — local engine semantics on a standalone radix-server +# ============================================================================= + + +def test_take_insert_match_recycle(env): + engine, _server = env.make("/cers_basic") + + seq = FakeSeq(block_hashes=_hashes(seed=1, num=4)) + # Initial match: nothing. + r = engine.match(seq) + assert r.num_matched_blocks == 0 + r.release() + + # take 4 slots and insert. + slots = engine.take(num_required_blocks=4) + assert len(slots) == 4 + engine.insert(seq, slots, num_insert_blocks=4) + + # Match should now hit all 4 blocks, all of them this node's slots. + r2 = engine.match(seq) + assert r2.num_matched_blocks == 4 + assert r2.local_slots.size == 4 + np.testing.assert_array_equal(np.sort(r2.local_slots), np.sort(slots)) + r2.release() + + # Recycle a fresh allocation; tree-attached slots are not affected. + free_slots = engine.take(num_required_blocks=2) + engine.recycle(free_slots) + + +def test_insert_publishes_immediately(env): + """There is no ready bit: being in the tree IS being servable. + + Insert runs after the transfer on this backend, so a matched block is + complete by construction — there is no flag to withhold a span with, and a + single insert() is the whole publication. + """ + engine, _server = env.make("/cers_unready") + + seq = FakeSeq(block_hashes=_hashes(seed=2, num=6)) + slots = engine.take(num_required_blocks=6) + engine.insert(seq, slots, num_insert_blocks=6) + + r = engine.match(seq) + assert r.num_matched_blocks == 6 + r.release() + + +def test_recycle_returns_staged_slots(env): + """Slots whose transfer never landed are only reachable through recycle(). + + They were never attached to the tree, so no query finds them and eviction + cannot reclaim them — without recycle() they are lost for the life of the + region. + """ + engine, _server = env.make("/cers_recycle") + + before = engine.num_free_blocks + slots = engine.take(num_required_blocks=5) + assert engine.num_free_blocks == before - 5 + engine.recycle(slots) + assert engine.num_free_blocks == before + + # And nothing was published on the way through. + seq = FakeSeq(block_hashes=_hashes(seed=22, num=5)) + r = engine.match(seq) + assert r.num_matched_blocks == 0 + r.release() + + +def test_eviction_reclaims_inserted(env): + """A published span is immediately LRU-evictable. + + insert() runs after the transfer, so the span it attaches has no reader and + takes no ref — nothing has to be released to make it reclaimable. + """ + engine, _server = env.make("/cers_evict", blocks=2000) + + seq = FakeSeq(block_hashes=_hashes(seed=4, num=1500)) + s1 = engine.take(num_required_blocks=1500) + engine.insert(seq, s1, num_insert_blocks=1500) + # insert() reports nothing, so check the span landed rather than let the + # eviction assert below pass on an empty tree. + published = engine.match(seq) + assert published.num_matched_blocks == 1500 + published.release() + assert engine.num_free_blocks == 500 + # Allocate enough new blocks that eviction is forced (need > current free 500). + s2 = engine.take(num_required_blocks=1500) + assert len(s2) > 500 + + +def test_pinned_match_survives_eviction_pressure(env): + """The query pin (`lock=True`) is what keeps a matched prefix out of the + evictor's reach until `release()`.""" + engine, _server = env.make("/cers_pin", blocks=2000) + seq = FakeSeq(block_hashes=_hashes(seed=5, num=1500)) + engine.insert(seq, engine.take(1500), num_insert_blocks=1500) + + pinned = engine.match(seq) + assert pinned.num_matched_blocks == 1500 + # 500 free; everything else is pinned, so the take comes up short. + short = engine.take(num_required_blocks=1500) + assert len(short) == 500 + engine.recycle(short) + pinned.release() + evicting = engine.take(num_required_blocks=1500) + assert len(evicting) == 1500 + engine.recycle(evicting) + + +def test_standalone_region_has_no_peer(env): + """A single-node region: peer reuse is off and prefetch has nothing to pull.""" + engine, _server = env.make("/cers_local_only", peer_enabled=True) + seq = FakeSeq(block_hashes=_hashes(seed=12, num=3)) + slots = engine.take(num_required_blocks=3) + engine.insert(seq, slots, num_insert_blocks=3) + + assert engine.is_distributed is False + assert engine.peer_enabled is False # asked for, but world_size == 1 + assert bootstrap.radix_cluster_rank(engine.client) == 0 + assert engine.prefetch(seq) is None + result = engine.match(seq) + assert result.num_matched_blocks == 3 + assert result.finalize is not None + result.release() + # release() is what drops the query's pin, and it is idempotent. + assert result.finalize is None + result.release() + + +def test_local_range_intersects_the_window(): + """`local_range` bounds the range on BOTH sides, by slicing alone. + + This is the accessor the GET and PUT planners lean on instead of clamping + the match end themselves, so the contract is that a hit stopping short of + the window contributes nothing and one running past the window end is + trimmed to it. + """ + match = ShmRadixMatch( + num_matched_blocks=4, + local_slots=np.arange(40, 44, dtype=np.int64), + ) + # Wholly inside the hit. + assert match.local_range(1, 3).tolist() == [41, 42] + # Hit runs PAST the window end -> trimmed to the window. + assert match.local_range(0, 2).tolist() == [40, 41] + # Window runs past the hit -> trimmed to the hit, no error. + assert match.local_range(2, 99).tolist() == [42, 43] + # Hit stops short of the window start -> nothing of it is ours. + assert match.local_range(4, 9).tolist() == [] + assert match.local_range(7, 9).tolist() == [] + # Empty and inverted windows name no block; full window is the whole hit. + assert match.local_range(2, 2).tolist() == [] + assert match.local_range(3, 1).tolist() == [] + assert match.local_range(0, 4).tolist() == [40, 41, 42, 43] + + +SWA_W = 8 # == flexkv.common.config.RADIX_SWA_WINDOW_BLOCKS, literal on purpose: + # a drive-by change to the constant should fail here, visibly. +JOINT_MASK = (_engine_mod.COMPONENT_MASK_FULL | + _engine_mod.COMPONENT_MASK_SWA) + + +def _make_swa_engine(env, name: str, blocks: int = 2000, swa_slots: int = 64, + tokens_per_block: int = 16, window_blocks: int = SWA_W): + """A single region carrying the SWA component, and an engine that knows it.""" + from flexkv.common.config import SWAPoolConfig + cfg = _server_config(name, blocks, tokens_per_block, + swa_slots=swa_slots, window_blocks=window_blocks) + server = env.server(cfg) + engine = env.engine( + name, num_total_blocks=blocks, tokens_per_block=tokens_per_block, + swa_config=SWAPoolConfig(enabled=True, num_slots=swa_slots, + window_blocks=window_blocks)) + return engine, server + + +def _publish_full(engine, seq, num_blocks: int) -> np.ndarray: + slots = engine.take(num_blocks) + engine.insert(seq, slots, num_insert_blocks=num_blocks) + return slots + + +def _publish_swa(engine, seq, path_end: int, + window_blocks: int = SWA_W) -> np.ndarray: + k = min(path_end, window_blocks) + slots = engine.take(k, component=_SWA) + assert len(slots) == k, "SWA pool unexpectedly short in test setup" + engine.insert(seq, slots, num_insert_blocks=path_end, component=_SWA) + return slots + + +def test_swa_window_invisible_until_published_then_joint_hit(env): + """With `common_hit=20` the Full slots cover [0, 20) and the SWA slots cover + [12, 20) -- and before insert(SWA), the joint query matches NOTHING even + though Full alone matches 20.""" + engine, _server = _make_swa_engine(env, "/cers_swa_basic") + seq = FakeSeq(block_hashes=_hashes(41, 20), tokens_per_block=16) + + _publish_full(engine, seq, 20) + + match = engine.match(seq) + assert match.num_matched_blocks == 20 + assert len(match.swa_slots) == 0 and match.swa_start == 0 + match.release() + match = engine.match(seq, component_mask=JOINT_MASK) + assert match.num_matched_blocks == 0 + assert len(match.swa_slots) == 0 + match.release() + + swa_slots = _publish_swa(engine, seq, path_end=20) + + match = engine.match(seq, component_mask=JOINT_MASK) + assert match.num_matched_blocks == 20 # joint common hit + assert match.local_slots.size == 20 # Full covers [0, 20) + assert match.swa_start == 12 # max(0, 20 - 8) + assert len(match.swa_slots) == SWA_W # window covers [12, 20) + assert sorted(match.swa_slots.tolist()) == sorted(swa_slots.tolist()) + match.release() + + +def test_swa_short_path_window_starts_at_zero(env): + """A path shorter than W publishes a window over the whole path: k=n slots, + swa_start=0 -- the `k = min(path_end, W)` boundary.""" + engine, _server = _make_swa_engine(env, "/cers_swa_short") + seq = FakeSeq(block_hashes=_hashes(42, 5), tokens_per_block=16) + + _publish_full(engine, seq, 5) + _publish_swa(engine, seq, path_end=5) + + match = engine.match(seq, component_mask=JOINT_MASK) + assert match.num_matched_blocks == 5 + assert match.swa_start == 0 + assert len(match.swa_slots) == 5 + match.release() + + +def test_swa_take_is_all_or_none_and_the_query_pin_protects_the_window(env): + """`allocate_slots(k, SWA)` returns k slots or NOTHING. With the pool sized + to exactly one window, a joint match's pin keeps that window un-evictable + (empty take); releasing the pin frees it for eviction (full take).""" + engine, _server = _make_swa_engine(env, "/cers_swa_allornone", swa_slots=SWA_W) + seq = FakeSeq(block_hashes=_hashes(43, 20), tokens_per_block=16) + + _publish_full(engine, seq, 20) + _publish_swa(engine, seq, path_end=20) + + match = engine.match(seq, component_mask=JOINT_MASK) + assert len(match.swa_slots) == SWA_W + empty = engine.take(SWA_W, component=_SWA) + assert len(empty) == 0 # all pinned -> all or none + match.release() + + evicted = engine.take(SWA_W, component=_SWA) + assert len(evicted) == SWA_W # pin gone -> window evictable + engine.recycle(evicted, component=_SWA) + + +def test_swa_insert_without_full_path_is_benign_and_recycles(env): + """FULL_PATH_MISSING (Full path evicted/absent under a pending SWA publish) + must cost the window, not the task: insert() warns, radixshmem auto-recycles + the whole batch, and the pool is whole again.""" + engine, _server = _make_swa_engine(env, "/cers_swa_orphan", swa_slots=SWA_W) + seq = FakeSeq(block_hashes=_hashes(44, 20), tokens_per_block=16) + + swa_slots = engine.take(SWA_W, component=_SWA) + assert len(swa_slots) == SWA_W + # No Full path published: refused with FULL_PATH_MISSING, not raised. + engine.insert(seq, swa_slots, num_insert_blocks=20, component=_SWA) + + again = engine.take(SWA_W, component=_SWA) + assert len(again) == SWA_W # auto-recycled, none leaked + engine.recycle(again, component=_SWA) + + +def test_swa_window_blocks_one_stores_a_single_slot_window(env): + """W comes from the region's config, not a constant: window_blocks=1 (the + SGLang DSv4 shape, window inside one page) publishes one-slot windows.""" + engine, _server = _make_swa_engine(env, "/cers_swa_w1", window_blocks=1) + seq = FakeSeq(block_hashes=_hashes(45, 20), tokens_per_block=16) + + _publish_full(engine, seq, 20) + _publish_swa(engine, seq, path_end=20, window_blocks=1) + + match = engine.match(seq, component_mask=JOINT_MASK) + assert match.num_matched_blocks == 20 + assert match.swa_start == 19 # max(0, 20 - 1) + assert len(match.swa_slots) == 1 + match.release() + + +# ============================================================================= +# Part 1b — the data plane: SlotStore as the CPU pool, geometry, server process +# ============================================================================= + + +def test_slot_store_pool_is_the_cpu_pool(env): + """The FULL pool of the server's SlotStore is FlexKV's CPU buffer: slot id + == block index, stride == block bytes, and a second attach by name (what a + transfer worker does) sees the same bytes.""" + torch = pytest.importorskip("torch") + from flexkv.storage.allocator import SlotStoreTensorHandle, slot_store_pool_tensor + + engine, _server = env.make("/cers_store", blocks=64) + store = engine.client.store + pool = store.pool(FULL) + assert int(pool.num_slots) == 64 + assert int(pool.slot_bytes) == SLOT_BYTES # exact stride, no padding + + slots = engine.take(3) + for i, slot in enumerate(slots): + engine.client.slot_view(int(slot))[:] = bytes([i + 1]) * SLOT_BYTES + + tensor = slot_store_pool_tensor(store, FULL, torch.uint8, 64 * SLOT_BYTES) + assert tensor.shape == (64 * SLOT_BYTES,) + for i, slot in enumerate(slots): + block = tensor[int(slot) * SLOT_BYTES:(int(slot) + 1) * SLOT_BYTES] + assert block.unique().tolist() == [i + 1] + + # A typed view (fp16) over the same pool: 64 blocks x SLOT_BYTES/2 elements. + typed = slot_store_pool_tensor(store, FULL, torch.float16, 64 * SLOT_BYTES // 2) + assert typed.dtype == torch.float16 and typed.numel() == 64 * SLOT_BYTES // 2 + + handle = SlotStoreTensorHandle(data_name=store.name, + hugepage_path=engine.client.info.hugepage_path, + kind=int(FULL), num_elements=64 * SLOT_BYTES, + dtype=torch.uint8) + worker_view = handle.get_tensor() # re-attached by name + first = int(slots[0]) + assert worker_view[first * SLOT_BYTES:(first + 1) * SLOT_BYTES].unique().tolist() == [1] + # And writes through the worker's view are what the owner reads. + worker_view[first * SLOT_BYTES] = 200 + assert bytes(engine.client.slot_view(first)[:1]) == b"\xc8" + engine.recycle(slots) + + +def test_slot_align_keeps_the_stride_exact(): + """slot_align is the largest power of two <= 4096 dividing every pool's + slot bytes, so radixshmem's round-up leaves the stride == slot bytes.""" + align = bootstrap.slot_align_for + assert align(2359296) == 4096 # Qwen3-8B block: 2^18 x 9 + assert align(149760 * 61) == 256 # 9135360 = 2^8 x 35685 + assert align(12345) == 1 # odd -> byte stride + assert align(4096 * 7, 1024 * 3) == 1024 + assert align(0, 8192) == 4096 # absent pools do not constrain + + +def _configs(num_cpu_blocks: int = 64, swa_slots: int = 0): + torch = pytest.importorskip("torch") + from flexkv.common.config import CacheConfig, ModelConfig, SWAPoolConfig + model_config = ModelConfig(num_layers=2, num_kv_heads=4, head_size=64, + dtype=torch.float16, tp_size=1, dp_size=1) + cache_config = CacheConfig(tokens_per_block=16, enable_cpu=True, enable_ssd=False, + num_cpu_blocks=num_cpu_blocks) + if swa_slots: + cache_config.swa = SWAPoolConfig(enabled=True, num_slots=swa_slots, + num_swa_layers=1, bytes_per_token_per_layer=64, + window_blocks=SWA_W) + return model_config, cache_config + + +def test_expected_geometry_mirrors_the_storage_engine_layout(): + """One FULL slot is one CPU block exactly as StorageEngine lays it out: + 2 layers x 2 (K,V) x 16 tokens x 4 heads x 64 x fp16 = 32768 B; one SWA + slot is one SWA page: 1 layer x 16 tokens x 64 B.""" + model_config, cache_config = _configs(num_cpu_blocks=64, swa_slots=16) + geo = bootstrap.expected_geometry(model_config, cache_config) + assert geo.tokens_per_block == 16 + assert geo.full_slots == 64 and geo.full_slot_bytes == 32768 + assert geo.swa_slots == 16 and geo.swa_slot_bytes == 1024 and geo.swa_window_blocks == SWA_W + assert geo.slot_align == 1024 # gcd power of two of 32768 and 1024 + assert geo.data_bytes == 64 * 32768 + 16 * 1024 + + +def test_server_config_and_geometry_check(env): + """`build_radix_server_config` starts a server whose regions pass + `check_geometry`; a different expectation is rejected, not papered over.""" + model_config, cache_config = _configs(num_cpu_blocks=64, swa_slots=16) + rcfg = _radix_config(cluster_id=f"geo{os.getpid()}") + set_radixshmem_config(rcfg) + try: + cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) + assert cfg.index.name == bootstrap.radix_index_name(rcfg.local_id) + assert cfg.cluster.cluster_id == rcfg.cluster_id + # FlexKV's own defaults, where they differ from radixshmem's. + assert cfg.cluster.rht_slots_per_bucket == 4 + assert cfg.cluster.bootstrap_timeout_sec == 120 + assert cfg.index.data_pool_ratio == 8.0 + # one RHT registration chunk per 4096 tokens, whatever the block size + assert cfg.index.register_chunk_size == 4096 // cache_config.tokens_per_block + assert cfg.data.slot_align == 1024 + assert cfg.index.full_slots == 64 and cfg.index.swa_slots == 16 + env.server(cfg) + client = bootstrap.attach_radix_client(cfg.index.name, timeout_s=30) + try: + geo = bootstrap.expected_geometry(model_config, cache_config) + bootstrap.check_geometry(client, geo, "test") + assert int(client.store.pool(FULL).slot_bytes) == geo.full_slot_bytes + assert int(client.store.pool(_SWA).slot_bytes) == geo.swa_slot_bytes + assert bootstrap.radix_cluster_rank(client) == 0 + cache_config.num_cpu_blocks = 65 + with pytest.raises(ValueError, match="FULL slots"): + bootstrap.check_geometry( + client, bootstrap.expected_geometry(model_config, cache_config), "test") + finally: + client.close() + finally: + set_radixshmem_config(None) + + +def test_embedded_server_process_lifecycle(): + """The bootstrap DP process runs the radix-server as a spawned subprocess: + start() returns once it is ready, clients attach by name, shutdown() takes + the socket down with it.""" + model_config, cache_config = _configs(num_cpu_blocks=64) + rcfg = _radix_config(cluster_id=f"proc{os.getpid()}") + set_radixshmem_config(rcfg) + shm_radix_id = rcfg.local_id + cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) + _sweep_region(cfg.index.name, cfg.resolved_data_name) + server = bootstrap.RadixServerProcess(cfg) + try: + server.start(timeout_s=120) + assert server.cluster_rank == 0 + assert server.info["distributed"] is False + assert os.path.exists(bootstrap.radix_socket_path(shm_radix_id)) + client = bootstrap.attach_radix_client(cfg.index.name, timeout_s=30) + assert client.info.data_plane + assert int(client.mempool_total()) == 64 + client.close() + finally: + server.shutdown() + set_radixshmem_config(None) + assert server.process is None + assert not os.path.exists(bootstrap.radix_socket_path(shm_radix_id)) + + +# ----------------------------------------------------------------------------- +# Part 1c — the radixshmem-mode YAML (flexkv.common.radixshmem_config): pass- +# through sections validated against shmradix's dataclasses, FlexKV's own +# defaults, the per-node overrides, and the startup checks +# (docs/radixshmem/config_zh.md section 6). + + +def _write_yaml(tmp_path, text: str) -> str: + path = tmp_path / "radixshmem.yaml" + path.write_text(text) + return str(path) + + +def test_radix_config_defaults_and_passthrough(tmp_path): + cfg = load_radixshmem_config(None) + assert cfg.cluster_id == "flexkv" and cfg.local_id == "flexkv" + assert not cfg.distributed and cfg.endpoint == "" + # FlexKV's defaults where they differ from radixshmem's; nothing else is set. + assert cfg.cluster == {"cluster_id": "flexkv", "bootstrap_timeout_sec": 120, + "rht_slots_per_bucket": 4} + assert cfg.index == {"data_pool_ratio": 8.0} and cfg.data == {} and cfg.server == {} + assert cfg.attach_timeout_s == 180.0 + + path = _write_yaml(tmp_path, """ +cluster: + cluster_id: prod + expected_min_nodes: 3 + registry: etcd://10.0.0.1:2379 + rpc_interface: eth0 + index_dev: mlx5_0 + rht_transport: xrc + peer_index_transport: dc + num_rht_shards: 2 + rht_shard_holders: "0,2" +data: + transfer_devices: mlx5_1,mlx5_2 + prefault: false +index: + background_evict_ratio: 0.1 +server: + rpc_workers: 8 +client: + prefetch_timeout_ms: 1000 +""") + cfg = load_radixshmem_config(path) + assert cfg.path == path and cfg.distributed and cfg.expected_min_nodes == 3 + assert cfg.cluster["rht_shard_holders"] == [0, 2] + assert cfg.data == {"transfer_devices": ["mlx5_1", "mlx5_2"], "prefault": False} + assert cfg.index == {"data_pool_ratio": 8.0, "background_evict_ratio": 0.1} + assert cfg.server == {"rpc_workers": 8} + assert cfg.client.prefetch_timeout_ms == 1000 and cfg.client.max_outstanding == 256 + # Every section constructs its shmradix dataclass as is. + shmradix.ClusterConfig(**cfg.cluster) + shmradix.IndexConfig(**cfg.index) + shmradix.DataPlaneConfig(data_bytes=1, full_slot_bytes=1, **cfg.data) + + +def test_radix_config_per_node_overrides(tmp_path): + path = _write_yaml(tmp_path, """ +cluster: + cluster_id: prod + expected_min_nodes: 2 + registry: etcd://10.0.0.1:2379 + rpc_interface: eth0 +""") + cfg = load_radixshmem_config(path, node_name="r1", rpc_address="127.0.0.1") + assert cfg.node_name == "r1" and cfg.rpc_address == "127.0.0.1" + # The explicit address must not lose to the file's interface (radixshmem + # lets the interface win), and co-located nodes get distinct regions. + assert cfg.cluster["rpc_interface"] == "" + assert cfg.local_id == "prod_r1" + assert bootstrap.radix_index_name(cfg.local_id) == "/shmradix_prod_r1_cpu" + # Without the interface, the address alone satisfies the cluster check. + path = _write_yaml(tmp_path, "cluster:\n expected_min_nodes: 2\n registry: etcd://h:1\n") + with pytest.raises(RadixShmemConfigError, match="rpc_interface"): + load_radixshmem_config(path) + load_radixshmem_config(path, rpc_address="10.0.0.5") + + +@pytest.mark.parametrize("text, match", [ + ("cluster:\n node_name: n0\n", "per-node"), + ("cluster:\n rpc_address: 10.0.0.1\n", "per-node"), + ("index:\n full_slots: 5\n", "geometry is derived"), + ("data:\n slot_align: 4096\n", "geometry is derived"), + ("cluster:\n transport: xrc\n", "unknown key"), + ("server:\n cluster: {}\n", "unknown key"), + ("peers: {}\n", "unknown section"), + ("- a\n", "must be a mapping"), + ("cluster:\n expected_min_nodes: 2\n rpc_interface: eth0\n", "registry"), + ("cluster:\n expected_min_nodes: 2\n registry: etcd://h:1\n rpc_interface: eth0\n" + " num_rht_shards: 3\n", "num_rht_shards"), + ("cluster:\n rht_slots_per_bucket: 3\n", "rht_slots_per_bucket"), + ("cluster:\n rht_transport: rc\n", "rht_transport"), + ("cluster:\n peer_index_transport: tcp\n", "peer_index_transport"), + ("cluster:\n remote_op_transport: xrc\n", "remote_op_transport"), + ("client:\n prefetch_max_inflight: 256\n", "max_outstanding"), + ("client:\n timeout: 5\n", "unknown key"), +]) +def test_radix_config_rejects(tmp_path, text, match): + with pytest.raises(RadixShmemConfigError, match=match): + load_radixshmem_config(_write_yaml(tmp_path, text)) + + +def test_radix_config_env_singleton_reloads_on_change(tmp_path, monkeypatch): + """`get_radixshmem_config` follows GLOBAL_CONFIG_FROM_ENV: the path and the + two per-node overrides; a test-installed config wins until reverted.""" + from flexkv.common.radixshmem_config import get_radixshmem_config + set_radixshmem_config(None) + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", None) + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_node_name", "") + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_rpc_address", "") + assert get_radixshmem_config().cluster_id == "flexkv" + path = _write_yaml(tmp_path, "cluster:\n cluster_id: other\n") + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", path) + assert get_radixshmem_config().cluster_id == "other" + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_node_name", "n7") + assert get_radixshmem_config().local_id == "other_n7" + set_radixshmem_config(_radix_config(cluster_id="pinned")) + assert get_radixshmem_config().cluster_id == "pinned" + set_radixshmem_config(None) + assert get_radixshmem_config().local_id == "other_n7" + + +def test_server_config_takes_the_yaml_sections(tmp_path): + """`build_radix_server_config` passes the four sections through and keeps + the geometry / naming its own.""" + model_config, cache_config = _configs(num_cpu_blocks=64) + cfg_path = _write_yaml(tmp_path, """ +cluster: + cluster_id: yamlsrv + expected_min_nodes: 2 + registry: etcd://10.0.0.1:2379 + rpc_interface: eth0 + index_dev: mlx5_3 + rht_transport: dc + num_rht_shards: 1 +data: + transfer_devices: [mlx5_4] + prefault: false + max_pending_jobs: 7 +index: + data_pool_ratio: 5.5 + register_chunk_size: 64 +server: + rpc_workers: 3 + endpoint: unix:///dev/shm/yamlsrv.sock +""") + rcfg = load_radixshmem_config(cfg_path) + cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) + assert cfg.index.name == "/shmradix_yamlsrv_cpu" and cfg.index.full_slots == 64 + assert cfg.index.data_pool_ratio == 5.5 + assert cfg.index.register_chunk_size == 64 # the file wins over the derived default + assert cfg.resolved_data_name == "/shmradix_yamlsrv_cpu_data" + assert cfg.data.transfer_devices == ["mlx5_4"] and cfg.data.max_pending_jobs == 7 + assert cfg.data.prefault is False and cfg.data.full_slot_bytes == 32768 + assert cfg.cluster.expected_min_nodes == 2 and cfg.cluster.index_dev == "mlx5_3" + assert cfg.cluster.rht_transport == "dc" and cfg.cluster.peer_index_transport == "xrc" + assert cfg.cluster.rht_slots_per_bucket == 4 and cfg.cluster.node_name == "" + assert cfg.rpc_workers == 3 and cfg.endpoint == "unix:///dev/shm/yamlsrv.sock" + assert cfg.distributed + + +# ============================================================================= +# Part 2 — planning on the radixshmem backend (`RadixShmemCacheEngine._plan_get`, +# `_plan_prefetch`, `_plan_put`, and the abort path of their handles), plus +# KVTaskEngine's handling of a job-backed prefetch task. Synthetic matches, no +# region. +# ============================================================================= + +TOKENS_PER_BLOCK = 16 + + +class FakeJob: + """Stand-in for `shmradix.PullJob`: what `_plan_prefetch` reads on return + (`local_hit`, `planned_hit`) and what `KVTaskEngine` polls.""" + + def __init__(self, local_hit: int, planned_hit: int, job_id: int = 7): + self.local_hit = local_hit + self.planned_hit = planned_hit + self.job_id = job_id + self.cancelled = False + self._result = None + + def done(self) -> bool: + return self._result is not None + + def wait(self, timeout=None): + if self.cancelled: + raise RuntimeError("cancelled") + if self._result is None: + raise TimeoutError("still running") + return self._result + + def cancel(self) -> None: + self.cancelled = True + + def complete(self, common_hit: int, remote_blocks: int, remote_bytes: int = 0, + source_rank: int = 1) -> None: + self._result = SimpleNamespace(common_hit=common_hit, remote_blocks=remote_blocks, + remote_bytes=remote_bytes, source_rank=source_rank, + finalize=lambda: None) + + +def _global_cache_engine(): + """Build a `RadixShmemCacheEngine` whose CPU tier is a plain `CacheEngineAccel`. + + The radixshmem planners run for real; only the tier is swapped (no region + needed) and its tree side is then stubbed by `_force_radixshmem`. + + Imported here rather than at module scope: `flexkv.cache.__init__` pulls in + `flexkv.c_ext` (libcudart), which Parts 1 and 3 deliberately avoid. + """ + try: + import torch + + from flexkv.cache.cache_engine import GlobalCacheEngine + from flexkv.cache.radix_shmem_planner import RadixShmemCacheEngine + from flexkv.common.config import CacheConfig, ModelConfig + except Exception as exc: # pragma: no cover - environment-dependent + pytest.skip(f"GlobalCacheEngine unavailable (needs CUDA + flexkv.c_ext): {exc}") + + class _PlannerOnAccelTier(RadixShmemCacheEngine): + def _build_cpu_cache_engine(self, cache_config, event_collector): + return GlobalCacheEngine._build_cpu_cache_engine( + self, cache_config, event_collector) + + model_config = ModelConfig( + num_layers=2, num_kv_heads=4, head_size=64, + dtype=torch.float16, tp_size=1, dp_size=1, + ) + cache_config = CacheConfig( + tokens_per_block=TOKENS_PER_BLOCK, + enable_cpu=True, enable_ssd=False, enable_remote=False, + num_cpu_blocks=256, + ) + return _PlannerOnAccelTier(cache_config, model_config) + + +def _local_match(slots, finalize=None) -> ShmRadixMatch: + slots = np.asarray(slots, dtype=np.int64) + return ShmRadixMatch(num_matched_blocks=len(slots), local_slots=slots, finalize=finalize) + + +def _force_radixshmem(engine, cpu_result: ShmRadixMatch, *, prefetch_job=None, + peer_enabled=None) -> None: + """Stub the tree side of a `_global_cache_engine()`. + + The tier keeps its real mempool (so `take` returns honest slot ids) but the + tree side is faked: the synthetic match names no real prefix, and the tier + here is a `CacheEngineAccel`, whose `insert` signature is a different one. + Records what a planner published (`engine.inserted_pools`) and what it + handed back (`engine.aborted_slots`), and the prefetch calls it made + (`engine.prefetch_calls`). + """ + engine._match_cpu = ( # type: ignore[method-assign] + lambda *args, **kwargs: cpu_result + ) + tier = engine.cpu_cache_engine + tier.peer_enabled = (prefetch_job is not None) if peer_enabled is None else peer_enabled + prefetch_calls = [] + + def _prefetch(sequence_meta, **kwargs): + prefetch_calls.append(kwargs) + return prefetch_job + + tier.prefetch = _prefetch # type: ignore[attr-defined] + inserted = [] + aborted = [] + + def _insert(sequence_meta, physical_block_ids, num_insert_blocks, + component=None, _sink=inserted): + _sink.append((num_insert_blocks, np.asarray(physical_block_ids))) + + def _recycle(physical_block_ids, component=None, + _orig=tier.recycle, _sink=aborted): + _sink.append(np.asarray(physical_block_ids)) + _orig(np.asarray(physical_block_ids)) + + tier.insert = _insert # type: ignore[method-assign] + tier.recycle = _recycle # type: ignore[method-assign] + engine.inserted_pools = inserted # type: ignore[attr-defined] + engine.aborted_slots = aborted # type: ignore[attr-defined] + engine.prefetch_calls = prefetch_calls # type: ignore[attr-defined] + + +def _fake_request(num_blocks: int, base: int = 0): + """(token_ids, token_mask, slot_mapping) for a fully-masked `num_blocks` window. + + `base` offsets the token ids so two requests name distinct sequences.""" + num_tokens = num_blocks * TOKENS_PER_BLOCK + token_ids = np.arange(base, base + num_tokens, dtype=np.int64) + token_mask = np.ones(num_tokens, dtype=np.bool_) + # GPU blocks 1000.. so they can't be confused with CPU slot ids. + slot_mapping = ( + np.repeat(np.arange(1000, 1000 + num_blocks), TOKENS_PER_BLOCK) + * TOKENS_PER_BLOCK + + np.tile(np.arange(TOKENS_PER_BLOCK), num_blocks) + ).astype(np.int64) + return token_ids, token_mask, slot_mapping + + +def _ops_by_type(graph): + ops = {} + for op in graph._op_map.values(): + ops.setdefault(op.transfer_type, []).append(op) + return ops + + +def _run_get(engine, num_blocks: int, cpu_result: ShmRadixMatch, *, prefetch=False, + prefetch_job=None, peer_enabled=None): + """Call get() through the radixshmem planners with a forced match result.""" + from flexkv.cache.cache_engine import DEFAULT_CACHE_STRATEGY + _force_radixshmem(engine, cpu_result, prefetch_job=prefetch_job, + peer_enabled=peer_enabled) + token_ids, token_mask, slot_mapping = _fake_request(num_blocks) + strategy = copy.deepcopy(DEFAULT_CACHE_STRATEGY) + if prefetch: + strategy.ignore_gpu = True + strategy.ignore_gds = True + graph, return_mask, callback, _op_cbs, _end = engine.get( + request_id=1, + token_ids=token_ids, + token_mask=token_mask, + slot_mapping=slot_mapping, + dp_client_id=0, + temp_cache_strategy=strategy, + ) + engine.get_callback = callback # type: ignore[attr-defined] + return graph, _ops_by_type(graph), return_mask + + +def test_local_hit_plans_one_h2d_and_releases_the_pin(): + """A local CPU hit is exactly one H2D read straight from the hit's slots; + the match pin lives until the graph completes.""" + engine = _global_cache_engine() + released = [] + cpu_slots = np.arange(40, 44, dtype=np.int64) + free_before = engine.cpu_cache_engine.mempool.num_free_blocks + graph, ops, return_mask = _run_get( + engine, 4, _local_match(cpu_slots, finalize=lambda: released.append(1))) + assert set(ops) == {TransferType.H2D} + h2d = ops[TransferType.H2D][0] + np.testing.assert_array_equal(h2d.src_block_ids, cpu_slots) + np.testing.assert_array_equal(h2d.dst_block_ids, np.arange(1000, 1004)) + assert return_mask.sum() == 4 * TOKENS_PER_BLOCK + # No staging taken, nothing to publish, nothing to give back. + assert engine.cpu_cache_engine.mempool.num_free_blocks == free_before + assert released == [] + engine.get_callback() # type: ignore[attr-defined] + assert released == [1] + assert engine.inserted_pools == [] # type: ignore[attr-defined] + assert engine.aborted_slots == [] # type: ignore[attr-defined] + + +def test_partial_hit_restores_the_prefix_only(): + """The hit ends inside the window: H2D covers the hit, the mask says so, + and nothing past it is planned (the miss is recomputed).""" + engine = _global_cache_engine() + cpu_slots = np.arange(20, 22, dtype=np.int64) + _graph, ops, return_mask = _run_get(engine, 4, _local_match(cpu_slots)) + h2d = ops[TransferType.H2D][0] + np.testing.assert_array_equal(h2d.src_block_ids, cpu_slots) + np.testing.assert_array_equal(h2d.dst_block_ids, np.arange(1000, 1002)) + assert bool(return_mask[:2 * TOKENS_PER_BLOCK].all()) + assert not bool(return_mask[2 * TOKENS_PER_BLOCK:].any()) + + +def test_miss_is_an_empty_plan_with_the_pin_dropped(): + engine = _global_cache_engine() + released = [] + graph, ops, return_mask = _run_get( + engine, 4, _local_match([], finalize=lambda: released.append(1))) + assert ops == {} + assert not bool(return_mask.any()) + assert released == [1] # dropped at plan time + engine.get_callback() # type: ignore[attr-defined] + assert engine.get_callback.prefetch_job is None # type: ignore[attr-defined] + + +def test_prefetch_starts_a_peer_pull(): + """A prefetch on a clustered tier is `RadixClient.pull_async`: the plan has no + ops, the job rides on the callback handle, and the mask is the planned pull + [local hit, planned hit).""" + engine = _global_cache_engine() + job = FakeJob(local_hit=1, planned_hit=4) + graph, ops, return_mask = _run_get( + engine, 4, _local_match(np.arange(20, 21)), prefetch=True, prefetch_job=job) + assert ops == {} + callback = engine.get_callback # type: ignore[attr-defined] + assert callback.prefetch_job is job + assert (callback.prefetch_local_hit_blocks, callback.prefetch_planned_hit_blocks) == (1, 4) + assert not bool(return_mask[:TOKENS_PER_BLOCK].any()) + assert bool(return_mask[TOKENS_PER_BLOCK:4 * TOKENS_PER_BLOCK].all()) + # Full|SWA mask only when the request is SWA-aware; plain prefetch is FULL. + (call,) = engine.prefetch_calls # type: ignore[attr-defined] + assert call["component_mask"] == _engine_mod.COMPONENT_MASK_FULL + assert call["query_end"] == 4 + assert call["timeout_ms"] == load_radixshmem_config(None).client.prefetch_timeout_ms + + +def test_prefetch_without_peers_is_an_empty_plan(): + engine = _global_cache_engine() + _graph, ops, return_mask = _run_get( + engine, 4, _local_match(np.arange(20, 22)), prefetch=True, peer_enabled=False) + assert ops == {} + assert not bool(return_mask.any()) + assert engine.get_callback.prefetch_job is None # type: ignore[attr-defined] + assert engine.prefetch_calls == [] # type: ignore[attr-defined] + + +def test_prefetch_backpressure_skips_the_peer_walk(): + """Too many pulls in flight: no pull_async, so the client never blocks.""" + engine = _global_cache_engine() + limit = load_radixshmem_config(None).client.prefetch_max_inflight + engine._prefetch_jobs = [FakeJob(0, 4) for _ in range(limit)] # none done + _graph, ops, return_mask = _run_get( + engine, 4, _local_match([]), prefetch=True, prefetch_job=FakeJob(0, 4)) + assert ops == {} and not bool(return_mask.any()) + assert engine.prefetch_calls == [] # type: ignore[attr-defined] + assert engine.get_callback.prefetch_job is None # type: ignore[attr-defined] + + +def _bare_task_engine(cache_config): + """A `KVTaskEngine` without transfer handles: enough of the task table for + the job polling paths (the same `__new__` trick the fallback managers use).""" + from flexkv.kvtask import KVTaskEngine + mgr = KVTaskEngine.__new__(KVTaskEngine) + mgr.cache_config = cache_config + mgr.tasks = {} + mgr.prefetch_jobs = {} + mgr.graph_to_task = {} + mgr.transfer_handles = [] + mgr.uncompleted_ops = {} + mgr.uncompleted_op_results = {} + mgr.uncompleted_graphs = {} + mgr.required_completed_count = 0 + return mgr + + +def _prefetch_task(task_id, job, num_blocks, local_hit, planned_hit): + from flexkv.common.transfer import TransferOpGraph + from flexkv.kvtask import KVTask, TaskStatus, TaskType + n = num_blocks * TOKENS_PER_BLOCK + return KVTask( + task_id=task_id, task_type=TaskType.PREFETCH, task_end_op_id=-1, + task_end_op_finished=False, status=TaskStatus.RUNNING, + token_ids=np.arange(n), slot_mapping=np.zeros(n, dtype=np.int64), + token_mask=np.ones(n, dtype=np.bool_), + graph=TransferOpGraph.create_empty_graph(), + return_mask=np.zeros(n, dtype=np.bool_), callback=None, op_callback_dict={}, + prefetch_job=job, prefetch_local_hit_blocks=local_hit, + prefetch_planned_hit_blocks=planned_hit) + + +def test_task_engine_completes_a_prefetch_from_its_job(): + """The PREFETCH task has an empty graph; `_update_tasks` polls the job and, + once done, reports the pulled range and completes the task.""" + try: + from flexkv.common.config import CacheConfig + from flexkv.kvtask import TaskStatus + except Exception as exc: # pragma: no cover + pytest.skip(f"kvtask unavailable (needs flexkv.c_ext): {exc}") + mgr = _bare_task_engine(CacheConfig(tokens_per_block=TOKENS_PER_BLOCK, num_cpu_blocks=64)) + + job = FakeJob(local_hit=1, planned_hit=4) + task = _prefetch_task(1, job, num_blocks=4, local_hit=1, planned_hit=4) + mgr.tasks[1] = task + mgr.prefetch_jobs[1] = job + mgr._process_empty_graph(1) # job pending: stays RUNNING + mgr._poll_prefetch_jobs() + assert task.status == TaskStatus.RUNNING and 1 in mgr.prefetch_jobs + + job.complete(common_hit=4, remote_blocks=3, remote_bytes=3 * SLOT_BYTES) + mgr._poll_prefetch_jobs() + assert task.status == TaskStatus.COMPLETED + assert 1 not in mgr.prefetch_jobs + assert not bool(task.return_mask[:TOKENS_PER_BLOCK].any()) + assert bool(task.return_mask[TOKENS_PER_BLOCK:].all()) + + # A shortfall (transfer refused, evicted before publish) narrows the mask. + job2 = FakeJob(local_hit=1, planned_hit=4) + task2 = _prefetch_task(2, job2, num_blocks=4, local_hit=1, planned_hit=4) + mgr.tasks[2] = task2 + mgr.prefetch_jobs[2] = job2 + job2.complete(common_hit=2, remote_blocks=1) + mgr._process_empty_graph(2) # done at first look + assert task2.status == TaskStatus.COMPLETED + assert bool(task2.return_mask[TOKENS_PER_BLOCK:2 * TOKENS_PER_BLOCK].all()) + assert not bool(task2.return_mask[2 * TOKENS_PER_BLOCK:].any()) + + # Nothing pulled at all: an empty mask, still a completed (not failed) task. + job3 = FakeJob(local_hit=1, planned_hit=4) + task3 = _prefetch_task(3, job3, num_blocks=4, local_hit=1, planned_hit=4) + mgr.tasks[3] = task3 + mgr.prefetch_jobs[3] = job3 + job3.complete(common_hit=1, remote_blocks=0) + mgr._poll_prefetch_jobs() + assert task3.status == TaskStatus.COMPLETED + assert not bool(task3.return_mask.any()) + + +def test_task_engine_cancel_hands_the_job_back(): + """Cancelling a job-backed prefetch cancels the job (the pull finishes in + the background) and never touches it again.""" + try: + from flexkv.common.config import CacheConfig + from flexkv.kvtask import TaskStatus + except Exception as exc: # pragma: no cover + pytest.skip(f"kvtask unavailable (needs flexkv.c_ext): {exc}") + mgr = _bare_task_engine(CacheConfig(tokens_per_block=TOKENS_PER_BLOCK, num_cpu_blocks=64)) + job = FakeJob(local_hit=0, planned_hit=4) + task = _prefetch_task(5, job, num_blocks=4, local_hit=0, planned_hit=4) + mgr.tasks[5] = task + mgr.prefetch_jobs[5] = job + mgr._cancel_task(5) + assert job.cancelled + assert task.status == TaskStatus.CANCELLED + assert 5 not in mgr.prefetch_jobs and 5 not in mgr.tasks + job.complete(common_hit=4, remote_blocks=4) + mgr._poll_prefetch_jobs() # nothing left to do + + +# ---- PUT planning ---- + +def _run_put(engine, num_blocks: int, cpu_result: ShmRadixMatch): + """Call put() through `_plan_put` with a forced match result.""" + _force_radixshmem(engine, cpu_result) + token_ids, token_mask, slot_mapping = _fake_request(num_blocks) + graph, return_mask, callback, _op_cbs, _end = engine.put( + request_id=2, + token_ids=token_ids, + token_mask=token_mask, + slot_mapping=slot_mapping, + dp_client_id=0, + ) + engine.put_callback = callback # type: ignore[attr-defined] + return graph, _ops_by_type(graph), return_mask + + +def test_put_with_no_match_stores_the_whole_window(): + """Cold start: nothing cached, so D2H covers every block and the span publishes.""" + engine = _global_cache_engine() + graph, ops, return_mask = _run_put(engine, num_blocks=4, + cpu_result=_local_match([])) + + op_d2h = ops[TransferType.D2H][0] + assert op_d2h.src_block_ids.size == 4 # nothing skipped + assert op_d2h.dst_block_ids.size == 4 + assert bool(return_mask.all()) + + engine.put_callback() # graph completion + # One insert, of the whole window, ending at block 4. + assert [(n, len(s)) for n, s in engine.inserted_pools] == [(4, 4)] + assert engine.aborted_slots == [] + + +def test_put_skips_the_cached_prefix(): + """A partial CPU match: D2H moves only the blocks past it, and they publish.""" + engine = _global_cache_engine() + released = [] + cached = np.arange(20, 23, dtype=np.int64) # 3 of 5 blocks already in CPU + graph, ops, return_mask = _run_put( + engine, num_blocks=5, + cpu_result=_local_match(cached, finalize=lambda: released.append(1))) + + op_d2h = ops[TransferType.D2H][0] + assert op_d2h.src_block_ids.size == 2 + # GPU blocks are 1000.., so the skipped prefix is visible in the ids read. + assert op_d2h.src_block_ids.tolist() == [1003, 1004] + # The staged slots are fresh — never the ones the match already holds. + assert not set(op_d2h.dst_block_ids.tolist()) & set(cached.tolist()) + # Only the newly stored blocks come back as stored. + assert not bool(return_mask[:3 * TOKENS_PER_BLOCK].any()) + assert bool(return_mask[3 * TOKENS_PER_BLOCK:].all()) + + assert released == [] # pinned while D2H runs + engine.put_callback() + # The span ends at block 5, and carries the 2 new blocks only. + assert [(n, len(s)) for n, s in engine.inserted_pools] == [(5, 2)] + assert engine.aborted_slots == [] + assert released == [1] # released after the publish + + +def test_put_with_fully_cached_window_does_nothing(): + """A match covering the window ends the PUT: nothing to store, nothing to arm.""" + engine = _global_cache_engine() + released = [] + graph, ops, return_mask = _run_put( + engine, num_blocks=3, + cpu_result=_local_match(np.arange(20, 23), finalize=lambda: released.append(1))) + assert ops == {} + assert not bool(return_mask.any()) + assert engine.inserted_pools == [] + assert engine.aborted_slots == [] + assert released == [1] + + +# ---- abort: the plan was cancelled before its graph launched ---- + +def test_get_abort_drops_the_pin(): + """A cancelled GET never runs its H2D; abort releases the match pin, and a + late completion on the consumed handle does nothing more.""" + engine = _global_cache_engine() + released = [] + _run_get(engine, 4, _local_match(np.arange(40, 44), + finalize=lambda: released.append(1))) + assert released == [] + engine.get_callback.abort() # type: ignore[attr-defined] + assert released == [1] + engine.get_callback() # type: ignore[attr-defined] + assert released == [1] + assert engine.inserted_pools == [] # type: ignore[attr-defined] + + +def test_put_abort_returns_the_staged_slots_and_drops_the_pin(): + """A cancelled PUT never runs its D2H: nothing is published, the staged + slots go back to the mempool and the match pin is released.""" + engine = _global_cache_engine() + released = [] + free_before = engine.cpu_cache_engine.mempool.num_free_blocks + _graph, ops, _mask = _run_put( + engine, num_blocks=5, + cpu_result=_local_match(np.arange(20, 23), finalize=lambda: released.append(1))) + staged = ops[TransferType.D2H][0].dst_block_ids + assert engine.cpu_cache_engine.mempool.num_free_blocks == free_before - 2 + assert released == [] + + engine.put_callback.abort() # type: ignore[attr-defined] + assert engine.inserted_pools == [] # type: ignore[attr-defined] + assert [s.tolist() for s in engine.aborted_slots] == [staged.tolist()] + assert engine.cpu_cache_engine.mempool.num_free_blocks == free_before + assert released == [1] + engine.put_callback() # type: ignore[attr-defined] + assert engine.inserted_pools == [] # consumed: no late publish + + +def test_prefetch_abort_leaves_the_job_alone(): + """Cancelling a job-backed prefetch is KVTaskEngine's business (it cancels + the job); the plan itself holds nothing to roll back.""" + engine = _global_cache_engine() + job = FakeJob(local_hit=0, planned_hit=4) + _run_get(engine, 4, _local_match([]), prefetch=True, prefetch_job=job) + engine.get_callback.abort() # type: ignore[attr-defined] + assert not job.cancelled + assert engine.aborted_slots == [] # type: ignore[attr-defined] + + +# ============================================================================= +# Part 2b — SWA planning on a real region (needs c_ext for GlobalCacheEngine) +# ============================================================================= + +SWA_ENV_BLOCKS = 64 + + +@contextlib.contextmanager +def _swa_global_engine(swa_slots: int = 2 * SWA_W, + num_blocks: int = SWA_ENV_BLOCKS, + window_blocks: int = SWA_W): + """A `RadixShmemCacheEngine` on a real radix-server with the SWA component. + + `cache_config.swa` + `enable_swa_transfer` turn on `swa_op_constructor`, and + the same config drives the bootstrap, so this also covers the + shm_radix_bootstrap side of the design. The pool is small on purpose -- + pin-release is asserted through exact take() counts. + """ + try: + import torch + from flexkv.cache.radix_shmem_planner import RadixShmemCacheEngine + except Exception as exc: # pragma: no cover - environment-dependent + pytest.skip(f"RadixShmemCacheEngine unavailable (needs CUDA + flexkv.c_ext): {exc}") + + from flexkv.common.config import CacheConfig, ModelConfig, SWAPoolConfig + + rcfg = _radix_config(cluster_id=f"swaplanner{os.getpid()}") + saved = {"enable_radixshmem": GLOBAL_CONFIG_FROM_ENV.enable_radixshmem} + GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True + set_radixshmem_config(rcfg) + + server = None + engine = None + try: + cache_config = CacheConfig( + tokens_per_block=TOKENS_PER_BLOCK, + enable_cpu=True, enable_ssd=False, enable_remote=False, + num_cpu_blocks=num_blocks, + ) + cache_config.swa = SWAPoolConfig(enabled=True, num_slots=swa_slots, + num_swa_layers=1, + bytes_per_token_per_layer=64, + window_blocks=window_blocks) + cache_config.enable_swa_transfer = True + model_config = ModelConfig(num_layers=2, num_kv_heads=4, head_size=64, + dtype=torch.float16, + tp_size=1, dp_size=1) + cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) + _sweep_region(cfg.index.name, cfg.resolved_data_name) + server = shmradix.RadixServer(cfg).start() + engine = RadixShmemCacheEngine(cache_config, model_config) + assert engine.swa_op_constructor.enabled, \ + "SWA gate should be on: enable_swa_transfer + radixshmem swa_enabled" + yield engine + finally: + if engine is not None and engine.cpu_cache_engine is not None: + engine.cpu_cache_engine.close() + if server is not None: + server.close() + for name, value in saved.items(): + setattr(GLOBAL_CONFIG_FROM_ENV, name, value) + set_radixshmem_config(None) + + +def _split_swa(ops_of_type): + full = [op for op in ops_of_type if not getattr(op, "is_swa", False)] + swa = [op for op in ops_of_type if getattr(op, "is_swa", False)] + return full, swa + + +def _real_seq(token_ids): + from flexkv.common.block import SequenceMeta + return SequenceMeta(token_ids=np.asarray(token_ids).copy(), + tokens_per_block=TOKENS_PER_BLOCK) + + +def test_put_then_get_swa_roundtrip_on_real_region(): + """PUT: the graph carries a 20-block Full D2H plus an 8-slot is_swa D2H, + both on the task-end barrier; before the completion callback a joint query + sees nothing; the callback publishes insert(FULL) then insert(SWA) and + releases the query. GET(swa_aware): one graph with a 20-block Full H2D plus + the 8-slot SWA H2D, both on the barrier; the query pin lives until the + callback and is gone after it.""" + with _swa_global_engine() as engine: + cpu = engine.cpu_cache_engine + num_total = SWA_ENV_BLOCKS + + # Record the publish order without breaking the real inserts. + published = [] + real_insert = cpu.insert + + def _recording_insert(*args, **kwargs): + published.append(kwargs.get("component")) + return real_insert(*args, **kwargs) + + cpu.insert = _recording_insert # type: ignore[method-assign] + + token_ids, token_mask, slot_mapping = _fake_request(20) + graph, put_mask, put_cb, _op_cbs, put_end = engine.put( + request_id=7, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0) + + full_d2h, swa_d2h = _split_swa(_ops_by_type(graph)[TransferType.D2H]) + assert len(full_d2h) == 1 and len(swa_d2h) == 1 + assert full_d2h[0].dst_block_ids.size == 20 + assert swa_d2h[0].src_block_ids.size == SWA_W + assert swa_d2h[0].dst_block_ids.size == SWA_W + put_end_preds = set(graph._op_map[put_end].predecessors) + assert {full_d2h[0].op_id, swa_d2h[0].op_id} <= put_end_preds + assert bool(put_mask.all()) + + pending = cpu.match(_real_seq(token_ids), component_mask=JOINT_MASK) + assert pending.num_matched_blocks == 0 + pending.release() + + put_cb() # graph completion + assert published == [shmradix.ComponentType.FULL, _SWA] + + after = cpu.match(_real_seq(token_ids), component_mask=JOINT_MASK) + assert after.num_matched_blocks == 20 + assert after.swa_start == 12 + assert len(after.swa_slots) == SWA_W + after.release() + + graph, get_mask, get_cb, _op_cbs, get_end = engine.get( + request_id=8, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0, swa_aware=True) + assert int(get_mask.sum()) == 20 * TOKENS_PER_BLOCK + + full_h2d, swa_h2d = _split_swa(_ops_by_type(graph)[TransferType.H2D]) + assert len(full_h2d) == 1 and len(swa_h2d) == 1 + assert full_h2d[0].src_block_ids.size == 20 + assert swa_h2d[0].src_block_ids.tolist() == after.swa_slots.tolist() + get_end_preds = set(graph._op_map[get_end].predecessors) + assert {full_h2d[0].op_id, swa_h2d[0].op_id} <= get_end_preds + + held = cpu.take(num_total) + assert len(held) == num_total - 20 + cpu.recycle(held) + + get_cb() # Full H2D + SWA H2D done + drained = cpu.take(num_total) + assert len(drained) == num_total + cpu.recycle(drained) + + +def test_put_degrades_to_full_only_when_the_swa_pool_is_exhausted(): + """An empty all-or-none SWA take drops the SWA leg -- no is_swa op, no SWA + staged insert -- and the Full plan proceeds untouched.""" + with _swa_global_engine(swa_slots=SWA_W) as engine: # exactly one window + cpu = engine.cpu_cache_engine + + tok_a, mask_a, sm_a = _fake_request(10) + _graph, _mask, put_cb_a, _cbs, _end = engine.put( + request_id=11, token_ids=tok_a, token_mask=mask_a, + slot_mapping=sm_a, dp_client_id=0) + put_cb_a() + pin = cpu.match(_real_seq(tok_a), component_mask=JOINT_MASK) + assert len(pin.swa_slots) == SWA_W + + tok_b, mask_b, sm_b = _fake_request(10, base=1_000_000) + graph, put_mask, put_cb_b, _cbs, _end = engine.put( + request_id=12, token_ids=tok_b, token_mask=mask_b, + slot_mapping=sm_b, dp_client_id=0) + full_d2h, swa_d2h = _split_swa(_ops_by_type(graph)[TransferType.D2H]) + assert len(full_d2h) == 1 and swa_d2h == [] + assert bool(put_mask.all()) + put_cb_b() + pin.release() + + full_only = cpu.match(_real_seq(tok_b)) + assert full_only.num_matched_blocks == 10 + full_only.release() + joint = cpu.match(_real_seq(tok_b), component_mask=JOINT_MASK) + assert joint.num_matched_blocks == 0 + joint.release() + + +def test_get_without_swa_aware_stays_full_only_on_swa_region(): + with _swa_global_engine() as engine: + token_ids, token_mask, slot_mapping = _fake_request(12) + _graph, _mask, put_cb, _cbs, _end = engine.put( + request_id=21, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0) + put_cb() + + graph, get_mask, get_cb, _cbs, _end = engine.get( + request_id=22, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0) + assert int(get_mask.sum()) == 12 * TOKENS_PER_BLOCK + full_h2d, swa_h2d = _split_swa(_ops_by_type(graph)[TransferType.H2D]) + assert len(full_h2d) == 1 and swa_h2d == [] + get_cb() + + +def _drive_put(engine, token_ids, token_mask, slot_mapping, request_id): + """put() + immediate completion; returns (graph, return_mask).""" + graph, return_mask, cb, _op_cbs, _end = engine.put( + request_id=request_id, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0) + cb() + return graph, return_mask + + +def test_reput_of_a_fully_cached_prefix_is_an_early_return(): + with _swa_global_engine() as engine: + tok, mask, sm = _fake_request(10) + _drive_put(engine, tok, mask, sm, request_id=41) + graph, return_mask = _drive_put(engine, tok, mask, sm, request_id=42) + assert _ops_by_type(graph) == {} + assert not bool(return_mask.any()) + + +def test_put_extension_releases_a_nonempty_match_pin_after_both_publishes(): + with _swa_global_engine() as engine: + cpu = engine.cpu_cache_engine + tok10, mask10, sm10 = _fake_request(10) + _drive_put(engine, tok10, mask10, sm10, request_id=51) + + tok20, mask20, sm20 = _fake_request(20) # same first 10 blocks + graph, return_mask, cb, _op_cbs, _end = engine.put( + request_id=52, token_ids=tok20, token_mask=mask20, + slot_mapping=sm20, dp_client_id=0) + full_d2h, swa_d2h = _split_swa(_ops_by_type(graph)[TransferType.D2H]) + assert full_d2h[0].dst_block_ids.size == 10 # only the extension moves + assert len(swa_d2h) == 1 # window rides along + held = cpu.take(SWA_ENV_BLOCKS) + assert len(held) == SWA_ENV_BLOCKS - 20 + cpu.recycle(held) + + cb() # FULL publish, SWA publish, release + drained = cpu.take(SWA_ENV_BLOCKS) + assert len(drained) == SWA_ENV_BLOCKS # pin gone, all evictable + cpu.recycle(drained) + + +def test_swa_get_of_a_shorter_prefix_misses(): + with _swa_global_engine() as engine: + tok, mask, sm = _fake_request(20) + _drive_put(engine, tok, mask, sm, request_id=61) + + short = 12 * TOKENS_PER_BLOCK + graph, get_mask, get_cb, _op_cbs, _end = engine.get( + request_id=62, token_ids=tok[:short], token_mask=mask[:short], + slot_mapping=sm[:short], dp_client_id=0, swa_aware=True) + assert int(get_mask.sum()) == 0 + assert _ops_by_type(graph) == {} + get_cb() + + graph, get_mask, get_cb, _op_cbs, _end = engine.get( + request_id=63, token_ids=tok[:short], token_mask=mask[:short], + slot_mapping=sm[:short], dp_client_id=0) + assert int(get_mask.sum()) == short + get_cb() + + +def test_short_path_put_and_get_use_k_smaller_than_w(): + with _swa_global_engine() as engine: + cpu = engine.cpu_cache_engine + tok, mask, sm = _fake_request(5) + graph, _ = _drive_put(engine, tok, mask, sm, request_id=71) + _full_d2h, swa_d2h = _split_swa(_ops_by_type(graph)[TransferType.D2H]) + assert swa_d2h[0].src_block_ids.size == 5 + + joint = cpu.match(_real_seq(tok), component_mask=JOINT_MASK) + assert (joint.num_matched_blocks, joint.swa_start, + len(joint.swa_slots)) == (5, 0, 5) + joint.release() + + graph, get_mask, get_cb, _op_cbs, _end = engine.get( + request_id=72, token_ids=tok, token_mask=mask, + slot_mapping=sm, dp_client_id=0, swa_aware=True) + assert int(get_mask.sum()) == 5 * TOKENS_PER_BLOCK + _full_h2d, swa_h2d = _split_swa(_ops_by_type(graph)[TransferType.H2D]) + assert swa_h2d[0].src_block_ids.size == 5 + get_cb() + + +def test_planner_uses_configured_window_blocks(): + with _swa_global_engine(window_blocks=1) as engine: + cpu = engine.cpu_cache_engine + tok, mask, sm = _fake_request(10) + graph, _ = _drive_put(engine, tok, mask, sm, request_id=81) + _full_d2h, swa_d2h = _split_swa(_ops_by_type(graph)[TransferType.D2H]) + assert swa_d2h[0].src_block_ids.size == 1 + + joint = cpu.match(_real_seq(tok), component_mask=JOINT_MASK) + assert (joint.num_matched_blocks, joint.swa_start, + len(joint.swa_slots)) == (10, 9, 1) + joint.release() + + graph, get_mask, get_cb, _op_cbs, _end = engine.get( + request_id=82, token_ids=tok, token_mask=mask, + slot_mapping=sm, dp_client_id=0, swa_aware=True) + assert int(get_mask.sum()) == 10 * TOKENS_PER_BLOCK + _full_h2d, swa_h2d = _split_swa(_ops_by_type(graph)[TransferType.H2D]) + assert swa_h2d[0].src_block_ids.size == 1 + get_cb() + + +# ============================================================================= +# Part 3 — opt-in two-node radixshmem cluster over RDMA: prefetch pulls a peer's +# blocks (index walk over RDMA, server-side RDMA READ of the SlotStore bytes), +# then a local match finds them and the bytes are the writer's. +# +# Two spawned processes each run a data-mode RadixServer (distinct data names +# and sockets on one host) and a `CacheEngineRadixShmem` attached to it. +# Gated behind FLEXKV_RUN_RADIX_PEER_TEST=1; needs an ACTIVE RDMA device, a +# shmradix built with RDMA + etcd + mooncake, and an etcd +# (FLEXKV_TEST_RADIX_REGISTRY, or `etcd` on PATH for a private one). +# ============================================================================= + +PEER_BLOCKS = 1170 +PEER_SLOT_BYTES = 65536 + + +def _active_rdma_devices() -> list: + """RDMA devices with an ACTIVE port, in FLEXKV_TEST_RDMA_DEVICES order.""" + def _has_active_port(device: str) -> bool: + for state in glob.glob(f"/sys/class/infiniband/{device}/ports/*/state"): + with contextlib.suppress(OSError): + with open(state) as f: + if "ACTIVE" in f.read(): + return True + return False + + requested = [d for d in os.getenv("FLEXKV_TEST_RDMA_DEVICES", "").split(",") if d] + candidates = requested or sorted( + os.path.basename(p) for p in glob.glob("/sys/class/infiniband/*")) + return [d for d in candidates if _has_active_port(d)] + + +def _free_port() -> int: + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +@pytest.fixture +def cluster(): + """(rdma device, etcd registry): skip when the RDMA prerequisites are + absent; start a private etcd when none is configured.""" + if os.getenv("FLEXKV_RUN_RADIX_PEER_TEST") != "1": + pytest.skip("set FLEXKV_RUN_RADIX_PEER_TEST=1 to run the RDMA test") + devices = _active_rdma_devices() + if not devices: + pytest.skip("no ACTIVE RDMA device found") + try: + from shmradix import _data + if not hasattr(_data, "DataPlaneRegistry"): + pytest.skip("shmradix built without etcd (no DataPlaneRegistry)") + except ImportError: + pytest.skip("shmradix built without the _data extension") + registry = os.getenv("FLEXKV_TEST_RADIX_REGISTRY", "") + proc = None + workdir = None + if not registry: + etcd = shutil.which("etcd") + if not etcd: + pytest.skip("set FLEXKV_TEST_RADIX_REGISTRY or put etcd on PATH") + client_port, peer_port = _free_port(), _free_port() + workdir = tempfile.mkdtemp(prefix="flexkv_radix_etcd_") + proc = subprocess.Popen( + [etcd, "--name", "t", "--data-dir", os.path.join(workdir, "data"), + "--listen-client-urls", f"http://127.0.0.1:{client_port}", + "--advertise-client-urls", f"http://127.0.0.1:{client_port}", + "--listen-peer-urls", f"http://127.0.0.1:{peer_port}", + "--initial-advertise-peer-urls", f"http://127.0.0.1:{peer_port}", + "--initial-cluster", f"t=http://127.0.0.1:{peer_port}"], + stdout=open(os.path.join(workdir, "etcd.log"), "w"), + stderr=subprocess.STDOUT) + registry = f"etcd://127.0.0.1:{client_port}" + deadline = time.monotonic() + 15 + while time.monotonic() < deadline: + with contextlib.suppress(OSError): + with socket.create_connection(("127.0.0.1", client_port), timeout=0.5): + break + time.sleep(0.2) + else: + proc.kill() + pytest.skip("private etcd did not come up") + try: + yield devices[0], registry + finally: + if proc is not None: + proc.terminate() + with contextlib.suppress(Exception): + proc.wait(10) + shutil.rmtree(workdir, ignore_errors=True) + + +def _peer_pattern(block: int, writer: int) -> bytes: + return bytes([(block * 7 + writer * 131 + 3) % 251 + 1]) * PEER_SLOT_BYTES + + +def _node_main(rank, prefix, cluster_id, registry, rdma_dev, ready, done, output, + local_head_blocks=0): + """One node: a data-mode RadixServer (in-process) plus the FlexKV engine.""" + try: + endpoint = f"unix:///dev/shm/{prefix.lstrip('/')}_r{rank}.sock" + data_name = f"{prefix}_data_r{rank}" + _sweep_region(prefix, data_name) + cluster_kwargs = dict( + expected_min_nodes=2, registry=registry, cluster_id=cluster_id, + node_name=f"r{rank}", rpc_address="0.0.0.0", index_dev=rdma_dev, + gid_idx=int(os.getenv("FLEXKV_TEST_RADIX_GID_IDX", "3")), + bootstrap_timeout_sec=60, rht_slots_per_bucket=4) + cfg = shmradix.RadixServerConfig( + index=shmradix.IndexConfig(name=prefix, tokens_per_block=16, + full_slots=PEER_BLOCKS), + data=shmradix.DataPlaneConfig( + data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, full_slot_bytes=PEER_SLOT_BYTES, + slot_align=4096, data_name=data_name, prefault=False, + transfer_devices=[rdma_dev]), + cluster=shmradix.ClusterConfig(**cluster_kwargs), + endpoint=endpoint, + ) + server = shmradix.RadixServer(cfg).start() # collective: waits for both + set_radixshmem_config(_radix_config().replace_server(endpoint=endpoint)) + engine = CacheEngineRadixShmem(prefix, num_total_blocks=PEER_BLOCKS, + tokens_per_block=16, peer_enabled=True) + if not engine.peer_enabled: + raise RuntimeError("engine did not see a distributed region") + cluster_rank = bootstrap.radix_cluster_rank(engine.client) + + hashes = np.arange(26, dtype=np.uint64) * 104729 + 101 + query_hashes = hashes[:-1] + + def _seq(block_hashes): + return FakeSeq(block_hashes=block_hashes.view(np.int64), tokens_per_block=16) + + if rank == 0: + sequence = _seq(hashes) + slots = engine.take(num_required_blocks=len(hashes)) + assert len(slots) == len(hashes) + for i, slot in enumerate(slots): + engine.client.slot_view(int(slot))[:] = _peer_pattern(i, writer=0) + # insert() publishes and, with peer_enabled, flushes the RHT so the + # reader can route to us. + engine.insert(sequence, slots, num_insert_blocks=len(hashes)) + output.put({"writer_rank": cluster_rank}) + ready.set() + if not done.wait(120): + raise TimeoutError("reader did not complete") + else: + if local_head_blocks > 0: + head = hashes[:local_head_blocks] + head_slots = engine.take(num_required_blocks=local_head_blocks) + assert len(head_slots) == local_head_blocks + for i, slot in enumerate(head_slots): + engine.client.slot_view(int(slot))[:] = _peer_pattern(i, writer=1) + engine.insert(_seq(head), head_slots, num_insert_blocks=local_head_blocks) + if not ready.wait(120): + raise TimeoutError("writer did not publish") + # The writer's RHT publication is asynchronous: prefetch until the + # pull brings the whole prefix home. + expect = len(query_hashes) + result = None + job = None + deadline = time.monotonic() + 30 + while time.monotonic() < deadline: + job = engine.prefetch(_seq(query_hashes), timeout_ms=20000) + result = job.wait(60) + if result.common_hit >= expect: + break + time.sleep(0.05) + if result is None or result.common_hit < expect: + raise AssertionError( + f"prefetch reached {result.common_hit if result else None} blocks, " + f"expected {expect}") + match = engine.match(_seq(query_hashes)) + bad = [] + for i, slot in enumerate(match.local_slots.tolist()): + writer = 1 if i < local_head_blocks else 0 + if bytes(engine.client.slot_view(int(slot))) != _peer_pattern(i, writer): + bad.append(i) + output.put({ + "num_matched": match.num_matched_blocks, + "job_local_hit": int(job.local_hit), + "job_planned_hit": int(job.planned_hit), + "remote_blocks": int(result.remote_blocks), + "remote_bytes": int(result.remote_bytes), + "source_rank": int(result.source_rank), + "bad_blocks": bad, + "reader_rank": cluster_rank, + }) + match.release() + done.set() + engine.close() + server.close() + except Exception: + output.put({"error": traceback.format_exc(), "rank": rank}) + ready.set() + done.set() + + +def _run_two_nodes(registry, rdma_dev, local_head_blocks=0): + ctx = mp.get_context("spawn") + ready = ctx.Event() + done = ctx.Event() + output = ctx.Queue() + prefix = f"/shmradix_peer_test_{os.getpid()}_{local_head_blocks}" + cluster_id = f"flexkv-peer-test-{os.getpid()}-{local_head_blocks}" + + processes = [ + ctx.Process(target=_node_main, + args=(rank, prefix, cluster_id, registry, rdma_dev, ready, done, + output, local_head_blocks)) + for rank in range(2) + ] + for process in processes: + process.start() + for process in processes: + process.join(timeout=240) + if process.is_alive(): + process.terminate() + process.join(timeout=5) + + messages = [] + while not output.empty(): + messages.append(output.get()) + errors = [message for message in messages if "error" in message] + assert not errors, errors + assert all(process.exitcode == 0 for process in processes) + return (next(m for m in messages if "num_matched" in m), + next(m for m in messages if "writer_rank" in m)) + + +def test_prefetch_pulls_a_peer_prefix_over_rdma(cluster): + """Node 1 holds nothing: the prefetch pulls all 25 blocks off node 0 and + the local match then serves them with node 0's bytes.""" + rdma_dev, registry = cluster + reader, writer = _run_two_nodes(registry, rdma_dev) + assert reader["job_local_hit"] == 0 + assert reader["job_planned_hit"] == 25 + assert reader["remote_blocks"] == 25 + assert reader["remote_bytes"] == 25 * PEER_SLOT_BYTES + assert reader["source_rank"] == writer["writer_rank"] + assert reader["num_matched"] == 25 + assert reader["bad_blocks"] == [] + + +def test_prefetch_extends_a_local_prefix_over_rdma(cluster): + """Node 1 holds blocks 0-9 itself: the prefetch pulls only 10-24, and the + match serves the head with node 1's bytes and the tail with node 0's.""" + rdma_dev, registry = cluster + reader, writer = _run_two_nodes(registry, rdma_dev, local_head_blocks=10) + assert reader["job_local_hit"] == 10 + assert reader["job_planned_hit"] == 25 + assert reader["remote_blocks"] == 15 + assert reader["source_rank"] == writer["writer_rank"] + assert reader["num_matched"] == 25 + assert reader["bad_blocks"] == [] + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-v"])) From f596705be1aaa8c07d32443401f802f1aeed22b9 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 15:52:04 +0800 Subject: [PATCH 08/21] shm_channel: carry the worker-measured durations over the result ring main #297 has every worker report wait_ms / xfer_ms / e2e_ms on the CompletedOp, and KVTaskEngine exports them as the flexkv_py_transfer_{wait,xfer,e2e}_duration_seconds histograms. The shm result ring packed a CompletedOp into a fixed 30 B record without them, so under radixshmem -- where every completion crosses the ring -- the three fields decoded as 0.0 and record_transfer_duration() skipped them as "never timed". The three histograms stayed empty on that path. Append the three values as f64 to the record (30 -> 54 B; still inside the 64 B result slot, and COMPLETED_OP_WIRE_SIZE / the slot-size assert follow the struct). block_results still does not cross the ring, as before. Test: the local round trip now asserts the durations survive, and a new case checks the record fits the default slot and keeps the failed flag alongside them. --- flexkv/transfer/shm_channel.py | 16 ++++++++++++---- tests/test_shm_channel.py | 24 ++++++++++++++++++++++-- 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/flexkv/transfer/shm_channel.py b/flexkv/transfer/shm_channel.py index 652a2defa..e1196c771 100644 --- a/flexkv/transfer/shm_channel.py +++ b/flexkv/transfer/shm_channel.py @@ -103,7 +103,7 @@ def _futex_wake(addr: int, count: int = 1) -> int: # Result ring holds one fixed-width CompletedOp record per slot (64 B × 65536 = 4 MB). DEFAULT_RESULT_SLOTS = 65536 # power of 2 -DEFAULT_RESULT_SLOT_SIZE = 64 # cache line; one 29 B CompletedOp record +DEFAULT_RESULT_SLOT_SIZE = 64 # cache line; one 54 B CompletedOp record _PAGE = 4096 @@ -120,12 +120,14 @@ def _round_up(x: int, m: int) -> int: # CompletedOp result-ring record: graph_id/op_id (i64), transfer_type (u8 index), -# num_blocks (u32), num_bytes (u64), flags (u8: bit0 = failed). +# num_blocks (u32), num_bytes (u64), flags (u8: bit0 = failed), then the +# worker-measured wait_ms / xfer_ms / e2e_ms (f64 each) so the CE side can +# export the same transfer-duration histograms as the in-process path. # block_results does NOT cross the ring: a failed graph degrades to whole-task # failure on the client side, which is the correct conservative reading for # every backend the shm path serves (mooncake's partial success never runs # through this channel). -_COMPLETED_OP = struct.Struct(" bytes: flags = 1 if getattr(op, "failed", False) else 0 return _COMPLETED_OP.pack( op.graph_id, op.op_id, tt_idx, op.num_blocks, op.num_bytes, flags, + float(getattr(op, "wait_ms", 0.0)), + float(getattr(op, "xfer_ms", 0.0)), + float(getattr(op, "e2e_ms", 0.0)), ) def decode_completed_op(buf: Any, off: int) -> Any: """Unpack a CompletedOp record from `buf` at byte offset `off`.""" from flexkv.common.transfer import CompletedOp - graph_id, op_id, tt_idx, num_blocks, num_bytes, flags = \ + graph_id, op_id, tt_idx, num_blocks, num_bytes, flags, wait_ms, xfer_ms, e2e_ms = \ _COMPLETED_OP.unpack_from(buf, off) tt = None if tt_idx == _TT_NONE else _TT_NAMES[tt_idx] return CompletedOp( @@ -160,6 +165,9 @@ def decode_completed_op(buf: Any, off: int) -> Any: transfer_type=tt, num_blocks=num_blocks, num_bytes=num_bytes, + wait_ms=wait_ms, + xfer_ms=xfer_ms, + e2e_ms=e2e_ms, failed=bool(flags & 1), ) diff --git a/tests/test_shm_channel.py b/tests/test_shm_channel.py index 7912dea03..5d4091856 100644 --- a/tests/test_shm_channel.py +++ b/tests/test_shm_channel.py @@ -133,15 +133,18 @@ def test_single_round_trip_local(): assert msgs == ["hello", {"k": 42}] # Result ring carries fixed-width CompletedOp records; all fields must - # round-trip, including the transfer_type string and the -1 sentinel. + # round-trip, including the transfer_type string, the worker-measured + # durations and the -1 sentinel. sent = [ CompletedOp(graph_id=7, op_id=3, transfer_type="H2D", - num_blocks=12, num_bytes=98304), + num_blocks=12, num_bytes=98304, + wait_ms=0.25, xfer_ms=1.5, e2e_ms=2.125), CompletedOp(graph_id=7, op_id=-1), # graph-completed sentinel ] ch.result_send(sent) out = ch.result_recv(timeout_s=0.0) assert out == sent + assert (out[0].wait_ms, out[0].xfer_ms, out[0].e2e_ms) == (0.25, 1.5, 2.125) assert out[1].is_graph_completed() finally: ch.close() @@ -150,6 +153,23 @@ def test_single_round_trip_local(): ctrl.unlink() +def test_result_record_carries_durations_and_fits_a_slot(): + """The fixed-width record must hold the #297 durations (f64) and still fit + the default 64 B result slot; a failed op keeps its flag alongside them.""" + from flexkv.transfer.shm_channel import ( + COMPLETED_OP_WIRE_SIZE, DEFAULT_RESULT_SLOT_SIZE, + decode_completed_op, encode_completed_op) + assert COMPLETED_OP_WIRE_SIZE <= DEFAULT_RESULT_SLOT_SIZE + op = CompletedOp(graph_id=1 << 40, op_id=(1 << 40) + 5, transfer_type="D2H", + num_blocks=3, num_bytes=3 * 4096, + wait_ms=12.75, xfer_ms=0.001, e2e_ms=1e6, failed=True) + rec = encode_completed_op(op) + assert len(rec) == COMPLETED_OP_WIRE_SIZE + back = decode_completed_op(memoryview(rec), 0) + assert back == op + assert back.is_graph_failed() is False and back.failed + + def test_submit_fragmentation(): """A payload larger than one slot must fragment and round-trip intact, interleaved with small single-slot messages.""" From fa509522c761e45e7f89838136cf112f55544f8b Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 21 Sep 2026 18:06:27 +0800 Subject: [PATCH 09/21] transfer_manager: the shm TE inherits CUDA_VISIBLE_DEVICES unless the parent is pinned narrower than it serves TransferManagerShmTEProcess cleared CUDA_VISIBLE_DEVICES for the TE subprocess whenever total_gpus > 1, so that a TE spawned by a DP rank pinned to one device (vLLM DP) could still cudaIpcOpenMemHandle from every rank's GPU. With sglang that clearing is wrong: sglang gives all TP workers one namespace (e.g. CUDA_VISIBLE_DEVICES=2,3 with --base-gpu-id 0), the workers register logical ids 0/1 inside it, and the TE, renumbered to the full set, opened those handles from physical GPUs 0 and 1 while the model ran on 2/3. Transfers only succeeded through NVLink peer access -- verified on H20-GPU-24 with DeepSeek-V4-Flash TP2 (flexkv-test/tp2-cvd run C). Decide from the namespace instead of the GPU count: keep the parent's CUDA_VISIBLE_DEVICES when it lists at least as many devices as this node's TE serves (instance_num x gpus_per_node), and clear it only when the parent sees fewer (the per-rank pinned layout, where the registered ids are physical). The single-GPU case keeps its old behaviour (inherit). Same rule main #287 applies to KVServer / TransferManagerOnRemote. tests/test_shm_te_cvd.py covers the policy table. --- flexkv/transfer_manager.py | 49 ++++++++++++++++++++++++++------------ tests/test_shm_te_cvd.py | 17 +++++++++++++ 2 files changed, 51 insertions(+), 15 deletions(-) create mode 100644 tests/test_shm_te_cvd.py diff --git a/flexkv/transfer_manager.py b/flexkv/transfer_manager.py index 9af8124a5..0bcee9aa3 100644 --- a/flexkv/transfer_manager.py +++ b/flexkv/transfer_manager.py @@ -1551,6 +1551,26 @@ def shutdown(self) -> None: flexkv_logger.info("TransferManagerMultiNodeHandle shutdown complete") +def shm_te_clears_cuda_visible_devices(cvd: Optional[str], gpus_needed: int) -> bool: + """Whether TransferManagerShmTEProcess drops CUDA_VISIBLE_DEVICES for the TE. + + The TE opens every registered GPU buffer with cudaIpcOpenMemHandle, so it + must number the devices exactly as the workers did when they registered: + + * one CUDA_VISIBLE_DEVICES for the whole engine (sglang, single-GPU vLLM): + the workers register logical ids inside that namespace, so the TE has to + INHERIT it. Clearing it renumbers the devices and puts the TE on the wrong + physical GPUs (or, on one GPU, fails with "device >= 0 && device < num_gpus"). + * one CUDA_VISIBLE_DEVICES per DP rank (vLLM DP, each rank pinned to its own + device): the parent sees fewer GPUs than this node's TE has to reach, and + the ids the workers registered are physical. Only then is it cleared. + """ + if cvd is None: + return False + visible = len([d for d in cvd.split(",") if d.strip()]) + return visible < gpus_needed + + class TransferManagerShmTEProcess: """Spawns the single TE subprocess for the multi-DP shm path. @@ -1581,21 +1601,20 @@ def start(self) -> None: if self.process is not None and self.process.is_alive(): return from flexkv.transfer.shm_channel_handle import te_shm_main - # CRITICAL: clear CUDA_VISIBLE_DEVICES in the TE subprocess so it can - # cudaIpcOpenMemHandle from ALL DPs' GPUs (the parent scheduler may have - # it restricted to its own DP rank's device). mp.Process(spawn) inherits - # env from parent unless we override. Save+restore around .start(). - # - # BUT: only do this when there is genuinely more than one GPU to span - # (multi-DP/TP/CP/PP). For a single-GPU deployment (total_gpus == 1, - # e.g. vLLM serve on one restricted GPU), clearing CVD renumbers devices - # in the TE subprocess so it no longer matches the device ordinal the - # worker recorded in its TensorSharedHandle — cudaIpcOpenMemHandle then - # fails with "device >= 0 && device < num_gpus". Leave CVD untouched so - # the TE and the registering worker agree on device numbering. - clear_cvd = self.model_config.total_gpus > 1 - _saved_cuda = (os.environ.pop("CUDA_VISIBLE_DEVICES", None) - if clear_cvd else None) + # mp.Process(spawn) inherits the parent's env. The TE keeps the parent's + # CUDA_VISIBLE_DEVICES whenever that namespace covers the GPUs it serves + # (the workers registered ids inside it); it is cleared only when the + # parent is pinned to fewer GPUs than the node's TE must reach. See + # shm_te_clears_cuda_visible_devices. Pop / restore around .start(). + cvd = os.environ.get("CUDA_VISIBLE_DEVICES") + gpus_needed = self.model_config.instance_num * self.model_config.gpus_per_node + clear_cvd = shm_te_clears_cuda_visible_devices(cvd, gpus_needed) + if cvd is not None: + flexkv_logger.info( + f"TransferManagerShmTEProcess: parent CUDA_VISIBLE_DEVICES={cvd!r}, " + f"TE serves {gpus_needed} GPU(s) on this node -> " + f"{'clearing it for the TE (per-rank pinned layout)' if clear_cvd else 'the TE inherits it'}") + _saved_cuda = os.environ.pop("CUDA_VISIBLE_DEVICES", None) if clear_cvd else None try: self.process = self.mp_ctx.Process( target=te_shm_main, diff --git a/tests/test_shm_te_cvd.py b/tests/test_shm_te_cvd.py new file mode 100644 index 000000000..ec5701e31 --- /dev/null +++ b/tests/test_shm_te_cvd.py @@ -0,0 +1,17 @@ +"""CUDA_VISIBLE_DEVICES policy of the radixshmem shm TE subprocess.""" +import pytest + +from flexkv.transfer_manager import shm_te_clears_cuda_visible_devices + + +@pytest.mark.parametrize("cvd, needed, clears", [ + (None, 2, False), # nothing set: the TE sees every GPU, ids are physical + ("2,3", 2, False), # sglang: one namespace for all TP workers -> inherit + ("0,1,2,3", 2, False), # a wider namespace than needed still covers the TE + ("5", 1, False), # single-GPU deployment on a restricted device -> inherit + ("3", 8, True), # vLLM DP: rank pinned to one device, TE serves 8 -> clear + ("", 2, True), # empty value hides every GPU; drop it + ("GPU-aaaa,GPU-bbbb", 2, False), # UUID form counts the same way +]) +def test_shm_te_cvd_policy(cvd, needed, clears): + assert shm_te_clears_cuda_visible_devices(cvd, needed) is clears From 0a2cbaf70a148dc8c459116d00187ad1efd86c26 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 11:41:50 +0800 Subject: [PATCH 10/21] radixshmem: attach to the operator's radix-server; FlexKV brings the geometry and adopts the slot counts radixshmem e5ce067 turned the radix-server into a process that starts with nothing model-specific (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`) and takes its slot shape from a client's Geometry; the slot counts are planned from the server's byte budget. FlexKV follows: it no longer creates, sizes or supervises a radix-server. * shm_radix_bootstrap: RadixServerProcess, build_radix_server_config and the name derivation are gone. RadixGeometry is FlexKV's side of the geometry (tokens per block, bytes per CPU block and SWA page, SWA window, slot alignment; no counts) with to_shmradix(); attach_radix_client(name, geometry=...) builds shmradix.RadixClient(name, Geometry), retries while the server is not reachable yet, then wait_ready(); a server serving another geometry or one whose budget holds no slot is refused with the cause. check_geometry compares block size, slot bytes, strides and the SWA window (not counts); adopt_geometry writes the server's counts into CacheConfig.num_cpu_blocks / swa.num_slots. * KVManager: every DP process attaches with the geometry and adopts the counts before anything sizes a pool; local dp 0 only spawns the shm TE. distributed_node_id is the server's rank. CacheEngineRadixShmem and the TE attach the same way (idempotent Configure), the TE also re-checks the regions. Peer reuse follows the server (world_size > 1) instead of a YAML flag. * radixshmem_config: the YAML is `server` (name, endpoint, ready_timeout_s) and `client`; the former cluster / data / index sections are refused with a pointer to the radix-server flags. FLEXKV_RADIX_SERVER_LAUNCH_MODE, FLEXKV_RADIX_NODE_NAME and FLEXKV_RADIX_RPC_ADDRESS are removed. TE channel names derive from the server name. * tests: the unit suite starts ServerConfig servers with a budget sized so the planned counts equal the test's, and brings the geometry through the engine; new cases cover adoption, the geometry refusal and a server that starts after the client. The e2e tests start `python -m shmradix.cli` themselves (per node in the two-node test) and attach through the new YAML. * docs/radixshmem/config_zh.md rewritten (server flags vs FlexKV YAML, geometry hand-off and adoption, multi-engine sharing, migration table); examples replaced by radixshmem.yaml + radix_server_{single,multi}_node.sh. The TE adopts the counts again after its own recompute_cache_block_counts (which sizes from cpu_cache_gb and undid the KVManager's adoption on DSv4, where layer_groups make the recompute effective). Verified on H20-GPU-24: 56 unit tests (in-process ServerConfig servers), the single-node e2e (dp 1 and 2 on real GPUs, radix-server started by the test), and under sglang with DeepSeek-V4-Flash TP2 the attach, geometry hand-off, count adoption (FULL 8605 -> 8610, SWA 1024 -> 1052 from `radix-server --data-bytes 32G --swa-ratio 0.5`), connector cluster query and TE attach; the request round trip after the TE fix is still to be run (the node's GPUs were taken by other jobs). --- CHANGELOG.md | 1 + docs/radixshmem/config_zh.md | 318 +++++------ .../radix_server_multi_node.sh | 17 + .../radix_server_single_node.sh | 11 + examples/radixshmem_configs/radixshmem.yaml | 19 + .../radixshmem_multi_node.yaml | 24 - .../radixshmem_single_node.yaml | 17 - flexkv/cache/radix_shmem_engine.py | 28 +- flexkv/cache/radix_shmem_planner.py | 20 +- flexkv/common/config.py | 15 +- flexkv/common/radixshmem_config.py | 352 ++++-------- flexkv/integration/sglang/connector.py | 9 +- flexkv/kvmanager.py | 102 ++-- flexkv/server/shm_radix_bootstrap.py | 509 ++++++++---------- flexkv/transfer_manager.py | 16 +- tests/radixshmem/radix_e2e_common.py | 39 ++ .../radixshmem/test_e2e_radix_prefetch_p2p.py | 64 ++- tests/radixshmem/test_e2e_radix_shmem.py | 33 +- tests/radixshmem/test_radix_shmem_engine.py | 398 ++++++-------- 19 files changed, 876 insertions(+), 1116 deletions(-) create mode 100755 examples/radixshmem_configs/radix_server_multi_node.sh create mode 100755 examples/radixshmem_configs/radix_server_single_node.sh create mode 100644 examples/radixshmem_configs/radixshmem.yaml delete mode 100644 examples/radixshmem_configs/radixshmem_multi_node.yaml delete mode 100644 examples/radixshmem_configs/radixshmem_single_node.yaml diff --git a/CHANGELOG.md b/CHANGELOG.md index 0e671ea41..cd03b5afc 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Feature Universal: +- radixshmem mode attaches to an operator-run `radix-server` and no longer creates one: the server is started per node with a name and a byte budget (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`), FlexKV's clients hand it the geometry (`shmradix.RadixClient(name, Geometry)`: tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), the server plans the slot counts from its budget and every FlexKV process adopts them into `CacheConfig.num_cpu_blocks` / `swa.num_slots` (`shm_radix_bootstrap.adopt_geometry`; `cpu_cache_gb` no longer sizes the CPU tier in this mode). `RadixServerProcess`, `build_radix_server_config`, `FLEXKV_RADIX_SERVER_LAUNCH_MODE`, `FLEXKV_RADIX_NODE_NAME` and `FLEXKV_RADIX_RPC_ADDRESS` are gone; the radixshmem YAML shrinks to `server` (`name`, `endpoint`, `ready_timeout_s`) and `client`, and rejects the former `cluster` / `data` / `index` sections (they are radix-server flags now; migration table in `docs/radixshmem/config_zh.md`). A server serving another geometry is refused (`GeometryMismatch`), so every engine on one server runs the same model, page size and SWA configuration. Requires radixshmem e5ce067 or later. - radixshmem mode now uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. The radix-server runs as a subprocess of the bootstrap DP (`FLEXKV_RADIX_SERVER_LAUNCH_MODE`). radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). See `docs/radixshmem_integration.md` and `docs/radixshmem_cross_node.md` - radixshmem mode is configured by one YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `cluster` / `data` / `index` / `server` sections pass through by key to radixshmem's `ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig` (validated against the installed dataclasses, geometry keys rejected), a `client` section holds the prefetch limits; `index.register_chunk_size` defaults to `4096 / tokens_per_block` blocks (one RHT registration chunk per 4096 tokens). The file is global: `cluster.cluster_id` is the only namespace (etcd keys and every shm / socket / TE channel name), node identity derives from `cluster.rpc_interface`. The `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone; `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` remain as per-node overrides for co-located nodes. Examples in `examples/radixshmem_configs/`, reference `docs/radixshmem/config_zh.md`. - radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles now roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not. Peer reuse in this mode follows the radixshmem YAML (`distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 309117c9b..181da46db 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -1,19 +1,20 @@ # radixshmem 模式配置参考 -本文列出 FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)时的全部配置项: -哪些走环境变量、哪些走 YAML、哪些由 FlexKV 自己推导而禁止手工设置。 +FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)时,配置分两处: -配置分三层: - -| 层 | 载体 | 内容 | +| 谁 | 载体 | 内容 | |---|---|---| -| 进程级开关 | 环境变量 `FLEXKV_RADIX_*` | 是否启用、YAML 路径、server 启动方式,以及两个仅供同机多节点测试的 per-node 覆盖 | -| 集群配置 | YAML,`FLEXKV_RADIXSHMEM_CONFIG_PATH` 指向 | 全局,所有节点逐字节相同;键名与 radixshmem 的 dataclass 字段一致 | -| 几何 | 由 `ModelConfig` / `CacheConfig` 推导 | slot 数、slot 字节数、对齐、shm 名;不可配置 | +| 运维 | `radix-server` 命令行,每节点一个进程 | 名字、SlotStore 字节预算及 SWA 占比、hugepage、传输引擎、集群成员(etcd、网卡、rank)、索引调优 | +| FlexKV | 环境变量 + 一个很小的 YAML(`FLEXKV_RADIXSHMEM_CONFIG_PATH`) | 是否启用、attach 哪个 server、等待多久、prefetch 限额 | +| FlexKV 推导 | `ModelConfig` / `CacheConfig` | 几何:每 block 的 token 数、一个 CPU block 的字节数、一个 SWA page 的字节数、SWA 窗口、slot 对齐 | + +FlexKV **不再创建 radix-server**。它只实例化 `shmradix.RadixClient`:第一个 client 把几何交给 server, +server 按自己的字节预算规划各池的 slot 数并发布;之后每个 FlexKV 进程 attach 时把这些 slot 数采纳到 +`CacheConfig`(`num_cpu_blocks`、`swa.num_slots`)。也就是说,在这个模式下 CPU 层的容量由 +`radix-server --data-bytes` 决定,`cpu_cache_gb` 只是占位。 -同一个值只有一个来源。YAML 里没写的键取本文列出的默认值;没有 YAML 时全部取默认值,即单机模式。 -示例文件在 `examples/radixshmem_configs/`:`radixshmem_single_node.yaml` 和 `radixshmem_multi_node.yaml`。 -实现在 `flexkv/common/radixshmem_config.py`。 +实现:`flexkv/common/radixshmem_config.py`(YAML)、`flexkv/server/shm_radix_bootstrap.py`(几何、attach、采纳)。 +radixshmem 侧的接口见 radixshmem 仓库 `python/README.md`。 --- @@ -21,252 +22,183 @@ | 变量 | 默认 | 说明 | |---|---|---| -| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担,KVServer 不启动,每个 DP 进程各建一个 KVTaskEngine 并 attach 共享的 radix 区域。在 `flexkv` 首次 import 前设置。 | -| `FLEXKV_RADIXSHMEM_CONFIG_PATH` | 空 | 第 2 节 YAML 的路径。为空时所有键取默认值。 | -| `FLEXKV_RADIX_SERVER_LAUNCH_MODE` | `embedded` | `embedded`:dp0 进程以子进程方式启动 radix-server;`external`:attach 运维已启动的 radix-server。 | -| `FLEXKV_RADIX_NODE_NAME` | 空 | per-node 覆盖,见 3.3。生产部署不设。 | -| `FLEXKV_RADIX_RPC_ADDRESS` | 空 | per-node 覆盖,见 3.3。设了就忽略 YAML 的 `cluster.rpc_interface`。生产部署不设。 | +| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担,KVServer 不启动,每个 DP 进程各建一个 KVTaskEngine 并 attach 同一个 radix-server。在 `flexkv` 首次 import 前设置。 | +| `FLEXKV_RADIXSHMEM_CONFIG_PATH` | 空 | 第 2 节 YAML 的路径。为空时全部取默认值:attach 本机 `radix-server --name /flexkv`。 | 另有两个 FlexKV 通用变量在该模式下有约束: - `FLEXKV_CPU_LAYOUT` 必须是 `BLOCKFIRST`。一个 SlotStore slot 就是一个连续的 CPU block,LAYERFIRST 给不出这个布局。 -- `FLEXKV_HUGETLBFS_DIR`(默认 `/mnt/hugepages`):`server.hugepage_path` 为空且 `CacheConfig.use_hugepage_cpu_buffer` 为真时,radix 区域建在这个 hugetlbfs 挂载点下。 +- `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID`:同一节点上多个推理引擎共享同一个 radix-server 和同一个 TE 时用来区分实例(第 4.4 节)。 该模式与 `enable_ssd`、`enable_remote` 互斥,启动时报错。`enable_p2p_cpu` / `enable_p2p_ssd` 也必须为 False: -跨节点复用由 radix-server 自己完成(etcd + RDMA),在 YAML 使集群成为分布式(`expected_min_nodes > 1` 或 -`num_rht_shards > 1`)时自动开启,不再经过 FlexKV 的 Redis P2P 路径。 +跨节点复用由 radix-server 自己完成(etcd + RDMA),在它以集群参数启动时自动开启,不经过 FlexKV 的 Redis P2P 路径。 + +已移除的变量:`FLEXKV_RADIX_SERVER_LAUNCH_MODE`(不再有嵌入式启动)、`FLEXKV_RADIX_NODE_NAME`、 +`FLEXKV_RADIX_RPC_ADDRESS`(节点身份是 `radix-server --node-name` / `--rpc-address` 的事)。 --- ## 2. YAML 字段 -五个段。`cluster` / `data` / `index` / `server` 四段的键按名字直接构造 radixshmem 的 -`ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig`,用 -`dataclasses.fields()` 校验:未知键报错,几何键(2.6)报错。`client` 段是 FlexKV 自己的参数。 +两个段,都可省略。 -### 2.1 `cluster`(`shmradix.ClusterConfig`) +### 2.1 `server`:attach 哪个 radix-server | 键 | 默认 | 说明 | |---|---|---| -| `cluster_id` | `flexkv` | 集群命名空间。既是 etcd 键前缀 `radix//...`,也派生本机全部 shm 和 socket 名(见 4)。同一台机器上跑多个 FlexKV 实例时用不同的 `cluster_id` 区分。 | -| `expected_min_nodes` | `0` | 集群开关。`> 1` 时进入 etcd + RDMA 模式:bootstrap 等到 etcd 里登记的节点数达到 `max(expected_min_nodes, num_rht_shards)` 且稳定 `settle_ms` 后分配 rank。`world_size` 是实际观察到的节点数,可以大于该值。 | -| `registry` | `etcd://127.0.0.1:2379` | etcd 地址,集群模式必填。格式 `etcd://host:port`,多个成员在 scheme 之后用逗号或分号分隔,scheme 只写一次:`etcd://10.0.0.1:2379,10.0.0.2:2379`。见 3.4。 | -| `rpc_interface` | 空 | 解析 bootstrap IP 的网卡名,如 `bond0`(南北向管理网卡)。集群模式下必填(除非用 `FLEXKV_RADIX_RPC_ADDRESS` 覆盖)。每个节点解析出自己的 IP,节点身份自动派生为 `node`。 | -| `rpc_port` | `0` | bootstrap / XRC 监听端口,0 由系统分配。 | -| `settle_ms` | `500` | 成员集合稳定多久后开始 bootstrap。 | -| `bootstrap_timeout_sec` | `120` | rendezvous 超时。FlexKV 的 attach 超时是该值加 60 秒。 | -| `index_dev` | 空 | index 控制面(RHT 面和 peer index 面)用的 HCA,空为第一个可用设备,通常是 `mlx5_0` 即东西向计算网卡。建议指定南北向管理网卡的 HCA(如 `mlx5_bond_0`):控制面只有小消息,计算网卡留给 KV 字节。KV 字节的 HCA 是 `data.transfer_devices`。 | -| `gid_idx` | `3` | RoCE GID 索引。 | -| `rht_transport` | `xrc` | client 到 RHT shard holder 的 QP 类型,`xrc` 或 `dc`。 | -| `peer_index_transport` | `xrc` | remote walk 读 peer index 的 QP 类型,`xrc` 或 `dc`。 | -| `remote_op_transport` | `zmq` | remote insert / query 控制面,`zmq` 或 `dc`。FlexKV 不开 remote op,该字段不生效。 | -| `num_rht_shards` | `0` | RHT 分片数。0 为每节点一片。设了必须不大于节点数。 | -| `rht_shard_holders` | `[]` | 持有分片的 rank 列表,空为 rank 0 到 `num_rht_shards - 1`。rank 由 rendezvous 后按 `node` 字典序分配。 | -| `rht_slots_per_bucket` | `4` | RHT 每 bucket 的 slot 数,取 1 / 2 / 4 / 8。1 是盲覆盖,会丢路由项。 | -| `enable_remote_insert` | `false` | 透传。 | -| `enable_remote_query` | `false` | 透传。 | -| `zmq_listen_port` | `0` | 透传。 | - -`rht_transport` / `peer_index_transport` / `remote_op_transport` / `num_rht_shards` / `rht_shard_holders` / -`rht_slots_per_bucket` 只有 rank 0 的值生效,bootstrap 时经 etcd `/config` 广播给其他节点。全局 YAML -下各节点值本来相同,这一规则只是多一层保险。 - -禁止出现:`node_name`、`rpc_address`。它们是 per-node 值,全局 YAML 放不下;需要时走 3.3 的环境变量。 - -### 2.2 `data`(`shmradix.DataPlaneConfig` 的非几何字段) - -| 键 | 默认 | 说明 | -|---|---|---| -| `transfer_devices` | `[]` | KV 字节传输(mooncake)用的 HCA 列表,空为 mooncake 发现的全部设备。 | -| `transfer_protocol` | `rdma` | `rdma` 或 `tcp`。 | -| `transfer_ip` | 空 | 数据面 IP,空为 rpc 地址。 | -| `transfer_port` | `0` | 0 由引擎选。 | -| `transfer_metadata` | `P2PHANDSHAKE` | mooncake 元数据服务。 | -| `prefault` | `true` | server 启动时 MAP_POPULATE 整个 SlotStore。D2H 延迟可预测,页由 server 进程的 NUMA 策略放置;大池会拉长启动时间。 | -| `max_inflight` | `256` | 传输引擎在飞 batch 数。 | -| `max_pending_jobs` | `4096` | 排队 + 运行 + 未领取的 job 上限,超过则 Submit 被拒。 | -| `job_ttl_s` | `60.0` | 未领取 job 的保留秒数。 | - -禁止出现:`data_bytes`、`full_slot_bytes`、`swa_slot_bytes`、`mamba_slot_bytes`、`slot_align`、`data_name`。 +| `name` | `/flexkv` | `radix-server --name`,即索引 shm 名。以 `/` 开头。也派生默认 socket 和 FlexKV 自己的 TE channel 前缀(第 5 节)。 | +| `endpoint` | 空 | gRPC 端点。空为 `unix:///dev/shm/.sock`;server 以 `--endpoint` 改成 TCP 或别的路径时这里写同一个值。 | +| `ready_timeout_s` | `600` | 一个 FlexKV 进程等 server **可达且 ready** 的总时长。覆盖运维晚起 server、SlotStore prefault、集群 rendezvous(server 的 `--bootstrap-timeout`)。超时报错并给出启动命令。 | -### 2.3 `index`(`shmradix.IndexConfig` 的非几何字段) +### 2.2 `client`:FlexKV 侧参数 | 键 | 默认 | 说明 | |---|---|---| -| `data_pool_ratio` | `8.0` | 索引 DataPool 大小系数:`full_slots × ratio × (12 或 16)` 字节。 | -| `background_evict_ratio` | `0.05` | 后台驱逐比例,0 关闭。 | -| `max_nodes` | `0` | radix 节点池容量,0 自动。 | -| `register_chunk_size` | `4096 / tokens_per_block` | RHT 注册粒度(block 数)。FlexKV 的默认让一段覆盖 4096 个 token,与 block 大小无关(radixshmem 自身默认 128 block)。 | +| `prefetch_timeout_ms` | `5000` | 一次 prefetch 拉取的服务端超时。到期后 job 以本地命中的部分完成。 | +| `prefetch_max_inflight` | `128` | 每个 DP 进程在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | +| `max_outstanding` | `256` | `RadixClient` 未领取 job 的上限。 | -禁止出现:`name`、`tokens_per_block`、`full_slots`、`swa_slots`、`swa_window_blocks`、`mamba_slots`、`evict_policy`。 +### 2.3 不再接受的段 -### 2.4 `server`(`shmradix.RadixServerConfig` 顶层) +旧格式的 `cluster` / `data` / `index` 段出现时直接报错:这些键现在都是 `radix-server` 的命令行参数(第 7 节有对照表)。 -| 键 | 默认 | 说明 | -|---|---|---| -| `endpoint` | 空 | gRPC 端点。空为 `unix:///dev/shm/.sock`。server 监听和 client attach 都用它。 | -| `rpc_workers` | `32` | gRPC 工作线程数。每个有在飞 job 的 client 占一个。 | -| `hugepage_path` | 空 | index 和 SlotStore 的 hugetlbfs 挂载点。空时按第 1 节 `FLEXKV_HUGETLBFS_DIR` 的规则决定。 | +示例: -### 2.5 `client`(FlexKV 侧,不传给 radixshmem) +```yaml +# /etc/flexkv/radixshmem.yaml +server: + name: /flexkv + ready_timeout_s: 900 +client: + prefetch_timeout_ms: 5000 + prefetch_max_inflight: 128 +``` -| 键 | 默认 | 说明 | -|---|---|---| -| `prefetch_timeout_ms` | `5000` | 一次 prefetch 拉取的服务端超时。到期后 job 以本地命中的部分完成。 | -| `prefetch_max_inflight` | `128` | 每个 DP 进程在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | -| `max_outstanding` | `256` | `RadixClient` 未领取 job 的上限。 | +--- -### 2.6 由 FlexKV 推导、禁止手工设置的字段 +## 3. 几何与 slot 数 + +FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `shmradix.Geometry`): | 字段 | 来源 | |---|---| -| `index.name` | `/shmradix__cpu` | -| `index.tokens_per_block` | `CacheConfig.tokens_per_block` | -| `index.full_slots` | `CacheConfig.num_cpu_blocks` | -| `index.swa_slots` / `swa_window_blocks` | `CacheConfig.swa.num_slots` / `window_blocks`,SWA 未开启为 0 | -| `data.full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 block 字节数 | -| `data.swa_slot_bytes` | 同上,SWA 池按 uint8 | -| `data.data_bytes` | `full_slots × full_slot_bytes + swa_slots × swa_slot_bytes` | -| `data.slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,保证 slot stride 等于 block 大小 | -| `data.data_name` | `_data` | - -这些值在 YAML 中出现时启动报错,防止和 `CacheConfig` 静默冲突。TE 进程 attach 后还会用 -`check_geometry` 复核 server 端区域和 FlexKV 自己的布局一致。 - ---- +| `block_size` | `CacheConfig.tokens_per_block`(sglang 的 page size) | +| `full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 CPU block 字节数(每 PP 段层数 × 节点内 KV head 数 × head_size × kv_dim × dtype × tokens_per_block) | +| `swa_slot_bytes` / `swa_window_blocks` | `CacheConfig.swa` 开启时:一个 SWA page 的字节数(uint8)与窗口块数;未开启则没有 SWA 池 | +| `slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,保证 SlotStore stride 等于 block 字节数 | -## 3. 示例 +server 收到几何后的规划(radixshmem 的规则):`swa_slots = floor(swa_ratio × data_bytes / swa_stride)`, +`full_slots = (data_bytes − SWA 占用) / full_stride`。任一池算出 0 个 slot、模型有 SWA 而 `--swa-ratio` 为 0, +都在 configure 时拒绝,FlexKV 报 `cannot serve FlexKV's geometry`。 -### 3.1 单机 +**采纳**:attach 成功后 `adopt_geometry` 把 `pools.full.num_slots` 写进 `CacheConfig.num_cpu_blocks`, +`pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,日志形如 +`adopted radix-server /flexkv's slot counts: FULL 8605 slots (cpu_cache_gb had given 1524), SWA 1024 slots`。 +之后 TE 的 StorageEngine、cache engine、指标都用采纳后的值。 -不设 `FLEXKV_RADIXSHMEM_CONFIG_PATH` 即可。等价于: - -```yaml -cluster: - cluster_id: flexkv - expected_min_nodes: 0 -``` +**校验**:每个 attach 方(KVManager、cache engine、TE)用 `check_geometry` 复核 server 发布的 +`block_size`、各池 `slot_bytes`、SlotStore stride、SWA 窗口与自己的布局一致,不一致报错退出,不会静默错位传输。 +slot 数不在校验范围内,它们是 server 的。 -无 etcd、无 RDMA 依赖。多个 DP 进程共享一个 radix-server 和一个 TE。 +**同一 server 上的多个 client** 必须带相同的几何:相同模型、page size、SWA 配置。第二个不同的几何被 server 以 +`GeometryMismatch` 拒绝,FlexKV 报 `already serves another geometry`。单机下 TP 不同的同一模型通常几何相同 +(节点内 KV head 数与 TP 无关),以 `check_geometry` 为准。 -### 3.2 多机(全局配置,所有节点同一文件) +--- -```yaml -# /etc/flexkv/radixshmem.yaml -cluster: - cluster_id: prod_a - expected_min_nodes: 4 - num_rht_shards: 4 - registry: etcd://10.0.0.1:2379 - rpc_interface: bond0 # 南北向网卡;每节点解析自己的 IP,身份为 node - index_dev: mlx5_bond_0 # index 内部 RDMA 的 HCA,南北向网卡 - gid_idx: 3 - rht_transport: xrc - peer_index_transport: dc - rht_slots_per_bucket: 4 - bootstrap_timeout_sec: 120 -data: - transfer_devices: [mlx5_1, mlx5_2] # KV 字节传输的 HCA - prefault: true -index: - data_pool_ratio: 8.0 -server: - rpc_workers: 32 -client: - prefetch_timeout_ms: 5000 - prefetch_max_inflight: 128 -``` +## 4. 启动方式 -每个节点: +### 4.1 单机 ```bash +# 运维,每节点一次;nohup / systemd 皆可。DSv4 这类有 SWA 池的模型给 --swa-ratio。 +radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 + +# 推理引擎侧 export FLEXKV_ENABLE_RADIXSHMEM=1 -export FLEXKV_RADIXSHMEM_CONFIG_PATH=/etc/flexkv/radixshmem.yaml export FLEXKV_CPU_LAYOUT=BLOCKFIRST +# 不设 FLEXKV_RADIXSHMEM_CONFIG_PATH 即 attach /flexkv ``` -节点身份、shm 名、rank 全部自动派生,文件里没有任何 per-node 内容。 +server 起来后打印 `Waiting for a client's geometry`;FlexKV 的第一个进程 attach 时把几何交过去,server 建区域后 ready, +所有进程的 `wait_ready` 返回。server 晚于引擎启动也可以:FlexKV 在 `ready_timeout_s` 内重试连接。 -### 3.3 同机多节点(测试) +### 4.2 多机(一个集群) -两个 radix-server 在一台机器上时,同一网卡解析出同一 IP,身份会撞。用两个 per-node 环境变量区分: +每个节点各起一个 server,用相同的 `--cluster-id` 和 `--registry`;`--rpc-interface`(或 `--rpc-address`)给对端拨入的 IP, +`--node-name` 空时自动为 `node`: ```bash -# 进程 A -FLEXKV_RADIX_NODE_NAME=r0 FLEXKV_RADIX_RPC_ADDRESS=127.0.0.1 ... -# 进程 B -FLEXKV_RADIX_NODE_NAME=r1 FLEXKV_RADIX_RPC_ADDRESS=127.0.0.1 ... +radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 \ + --expected-min-nodes 4 --num-rht-shards 4 --rht-slots 4 \ + --registry etcd://10.0.0.1:2379 --cluster-id prod_a \ + --rpc-interface bond0 --index-dev mlx5_bond_0 --gid-idx 3 \ + --transfer-dev mlx5_1 --transfer-dev mlx5_2 --bootstrap-timeout 600 ``` -`FLEXKV_RADIX_RPC_ADDRESS` 设置后 FlexKV 清掉 YAML 的 `rpc_interface`(radixshmem 规则是 interface 优先, -不清会被覆盖回去)。两个进程仍共用同一份 YAML。 +集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`)由第一个拿到几何的节点发布到 +etcd `radix//geometry/`,其余 `waiting` 的节点采纳;各节点的 slot 数可以不同(预算可以不同)。 +FlexKV 侧每个节点同一份 YAML 即可。`ready_timeout_s` 要不小于 `--bootstrap-timeout`。 -设置了 `FLEXKV_RADIX_NODE_NAME` 时,本机命名前缀从 `` 变为 `_`(第 4 节), -两个进程的 SlotStore、socket 和 TE channel 因此互不冲突。每个进程还要各给一个 `FLEXKV_SERVER_RECV_PORT`。 +### 4.3 同机多节点(测试) -### 3.4 `registry` 的填法 +两个 server 在一台机器上:不同的 `--name`(或相同 name 加 `--endpoint`、`--data-name` 区分)、不同的 `--node-name`、 +`--rpc-address 127.0.0.1`。两个 FlexKV 进程各用一份 YAML,`server.name` / `server.endpoint` 指向自己的 server, +并各给一个 `FLEXKV_SERVER_RECV_PORT`。 -radixshmem 把 `registry` 去掉第一个 `://` 之前的 scheme 后,余下部分按逗号或分号切成 endpoint 列表, -交给 etcd 的 clientv3。因此: +### 4.4 一节点多引擎共享一个 server -- 单成员:`etcd://10.0.0.1:2379`。 -- 多成员:`etcd://10.0.0.1:2379,10.0.0.2:2379,10.0.0.3:2379`。scheme 只写一次;写成 - `etcd://a:2379,etcd://b:2379` 会把第二个 `etcd://b:2379` 原样当作 endpoint 传下去,连接失败。 -- 只支持明文连接,没有 TLS 和用户名密码的配置入口。拨号超时固定 5 秒。 -- 一个 etcd 可以服务多个集群,键空间由 `cluster_id` 隔开(`radix//...`);索引 rendezvous 和 - 数据面登记(`data/`)都在同一个 etcd 里。 -- etcd 不只在启动时用:节点的 lease keep-alive、`/peers` watch 和数据面登记贯穿整个运行期,etcd 不可用会导致 - lease 过期、节点从集群视图中消失。生产环境用 3 成员 etcd,并把全部成员写进 `registry`。 -- 每个节点必须能访问 `registry` 里的地址;默认值 `127.0.0.1:2379` 只适用于所有节点在同一台机器上的测试。 -- 同一进程内 etcd 连接是全局单例,首个 `init` 的 endpoint 生效;FlexKV 里索引和数据面都用 `cluster.registry`, - 不会出现两个不同地址。 +两个独立的推理引擎(各自的 FlexKV、各自的 GPU)attach 同一个 radix-server,互相命中对方存的 KV: ---- +```bash +# 引擎 A # 引擎 B +FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=1 +``` -## 4. 命名派生 +同一份 YAML。`instance 0` 的 dp0 拉起本节点唯一的 TE(`channels = instance_num × dp_size`),TE 等到 +`instance_num × gpus_per_node` 张 GPU 都注册才 ready,所以两个引擎都要启动。两边模型 / page size / SWA 配置必须相同 +(第 3 节)。node-local DP(多机 DP attention)路径下 `local_dp_client_id` 不带 instance,多实例暂不支持。 -所有名字来自 `cluster.cluster_id`。记本机前缀 `local_id`:未设 `FLEXKV_RADIX_NODE_NAME` 时就是 -`cluster_id`,设了则是 `_`(`RadixShmemConfig.local_id`)。 +--- + +## 5. 命名派生 | 对象 | 名字 | |---|---| +| index shm | `--name`;集群模式下 radixshmem 追加 `_`,attach 方只需 `--name` | +| SlotStore shm | `_data`(`--data-name` 可改) | +| gRPC socket | `/dev/shm/.sock`(`--endpoint` 可改;YAML `server.endpoint` 跟着改) | | etcd 键空间 | `radix//...` | -| index shm | `/shmradix__cpu`;集群模式下 radixshmem 再追加 `_`,attach 方只需 base name | -| SlotStore shm | `/shmradix__cpu_data` | -| gRPC socket | `/dev/shm/shmradix__cpu.sock` | -| TE shm channel | FlexKV 内部 IPC 名,以 `local_id` 为前缀 | +| FlexKV TE channel / ctrl | `/dev/shm/flexkv_te_ch__`、`flexkv_te_ctrl_`,`te_server_id` = `name` 去掉开头的 `/`(`/` 换成 `_`) | --- -## 5. 三套传输的区分 +## 6. 启动时校验 + +FlexKV 加载 YAML 时报错的情况:出现 `cluster` / `data` / `index` 段;未知段或未知键;`server.name` 不以 `/` 开头或含空白; +`server.ready_timeout_s <= 0`;`client.prefetch_max_inflight >= client.max_outstanding`;`client.prefetch_timeout_ms <= 0`。 -| 配置 | 取值 | 链路 | HCA | -|---|---|---|---| -| `cluster.rht_transport` | xrc / dc | client 向 RHT shard holder 写路由项 | `cluster.index_dev` | -| `cluster.peer_index_transport` | xrc / dc | remote walk 时对 peer 节点 index 的单边 RDMA read | `cluster.index_dev` | -| `cluster.remote_op_transport` | zmq / dc | remote insert / query 控制面,FlexKV 不使用 | zmq 走 TCP | -| `data.transfer_protocol` + `data.transfer_devices` | rdma / tcp | 两节点 SlotStore 之间的 KV 字节搬运(mooncake),即 `pull_async` 的实际拉取 | `data.transfer_devices` | +`CacheConfig` 侧:`FLEXKV_CPU_LAYOUT != BLOCKFIRST`;打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`; +SWA 开启但 `window_blocks < 1`。 -xrc 对每个目标一条 QP;dc 用一个 DC initiator 对所有目标,QP 数 O(1),需要 mlx5。任一 index 面为 dc 时 -server 建 DCT,client 两个面共用一个 DCI。FlexKV 的 `get_match` 只查本地,不走前两条;`pull_async` -在服务端规划时走 RHT 和 remote walk,随后的字节搬运走 mooncake。 +attach 时:`ready_timeout_s` 内连不上 server 报 `no radix-server named ... reachable ...(start it with radix-server --name ...)`; +server 一直在等几何或配置失败报 `not ready within ...`(带 server 的 `mode` 和 `last_error`);几何冲突见第 3 节。 --- -## 6. 启动时校验 +## 7. 从旧版迁移 -FlexKV 在加载 YAML 时检查以下条件,不满足直接报错,不等到 rendezvous: - -- 未知键、2.6 的几何键、`cluster.node_name`、`cluster.rpc_address` 出现在 YAML。 -- `expected_min_nodes > 1` 时 `registry` 为空,或 `rpc_interface` 与 `FLEXKV_RADIX_RPC_ADDRESS` 都为空。 -- `FLEXKV_RADIX_RPC_ADDRESS=0.0.0.0`:所有节点会派生出同一个身份。 -- `rht_shard_holders` 和 `transfer_devices` 接受列表或逗号分隔字符串,其他类型报错。 -- `num_rht_shards > expected_min_nodes`(`expected_min_nodes` 非 0 时)。 -- `rht_slots_per_bucket` 不在 {1, 2, 4, 8}。 -- `rht_transport` / `peer_index_transport` 不在 {xrc, dc};`remote_op_transport` 不在 {zmq, dc}。 -- `client.prefetch_max_inflight >= client.max_outstanding`。 -- `FLEXKV_CPU_LAYOUT != BLOCKFIRST`,或 `num_cpu_blocks <= 0`,或 SWA 池装不下一个 window。 -- `CacheConfig` 打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`。 - -TE attach 后 `check_geometry` 复核 server 端的 `tokens_per_block`、各池 slot 数、slot 字节数和 stride, -不一致则报错退出,不会静默错位传输。 +| 旧 YAML 键 | 现在 | +|---|---| +| `cluster.cluster_id` | `radix-server --cluster-id`(同时不再派生 shm 名;shm 名是 `--name`) | +| `cluster.expected_min_nodes` / `num_rht_shards` / `rht_shard_holders` / `rht_slots_per_bucket` | `--expected-min-nodes` / `--num-rht-shards` / `--rht-shard-holders` / `--rht-slots` | +| `cluster.registry` / `rpc_interface` / `rpc_port` / `settle_ms` / `bootstrap_timeout_sec` | `--registry` / `--rpc-interface` / `--rpc-port` / `--settle-ms` / `--bootstrap-timeout` | +| `cluster.index_dev` / `gid_idx` / `rht_transport` / `peer_index_transport` / `remote_op_transport` / `zmq_listen_port` | 同名 `--index-dev` 等 | +| `data.transfer_devices` / `transfer_protocol` / `transfer_ip` / `transfer_port` / `transfer_metadata` | `--transfer-dev`(可重复)/ `--transfer-protocol` / `--transfer-ip` / `--transfer-port` / `--transfer-metadata` | +| `data.prefault` / `max_inflight` / `max_pending_jobs` / `job_ttl_s` | `--no-prefault` / `--max-inflight` / `--max-pending-jobs` / `--job-ttl` | +| `index.data_pool_ratio` / `background_evict_ratio` / `max_nodes` / `register_chunk_size` | `--data-pool-ratio` / `--background-evict-ratio` / `--max-nodes` / `--register-chunk-tokens`(按 token 数) | +| `server.endpoint` | 保留:server 的 `--endpoint` 与 YAML `server.endpoint` 各写一次 | +| `server.rpc_workers` / `hugepage_path` | `--rpc-workers` / `--hugepage-path` | +| (由 FlexKV 推导的 slot 数、`data_bytes`) | slot 数由 `--data-bytes` 和 `--swa-ratio` 决定,FlexKV 采纳 | +| `FLEXKV_RADIX_SERVER_LAUNCH_MODE` | 移除,只有外部 server | +| `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` | `--node-name` / `--rpc-address` | diff --git a/examples/radixshmem_configs/radix_server_multi_node.sh b/examples/radixshmem_configs/radix_server_multi_node.sh new file mode 100755 index 000000000..a0894e881 --- /dev/null +++ b/examples/radixshmem_configs/radix_server_multi_node.sh @@ -0,0 +1,17 @@ +#!/bin/bash +# The operator's radix-server on ONE node of a cluster; run it on every node with +# the same --cluster-id / --registry. Peers dial the IP resolved from +# --rpc-interface (node identity defaults to node); the index control plane +# uses --index-dev, the KV bytes move over --transfer-dev. The first node that +# receives a geometry publishes it to etcd, the others adopt it; slot counts +# may differ per node (different budgets). FlexKV's ready_timeout_s must cover +# --bootstrap-timeout. +set -eu +radix-server --name "${NAME:-/flexkv}" \ + --data-bytes "${DATA_BYTES:-64G}" \ + ${SWA_RATIO:+--swa-ratio "$SWA_RATIO"} \ + --expected-min-nodes "${NODES:-2}" --num-rht-shards "${NODES:-2}" --rht-slots 4 \ + --registry "${REGISTRY:-etcd://10.0.0.1:2379}" --cluster-id "${CLUSTER_ID:-flexkv_prod}" \ + --rpc-interface "${RPC_INTERFACE:-bond0}" --index-dev "${INDEX_DEV:-mlx5_bond_0}" --gid-idx "${GID_IDX:-3}" \ + --transfer-dev "${TRANSFER_DEV:-mlx5_0}" \ + --bootstrap-timeout "${BOOTSTRAP_TIMEOUT:-600}" diff --git a/examples/radixshmem_configs/radix_server_single_node.sh b/examples/radixshmem_configs/radix_server_single_node.sh new file mode 100755 index 000000000..4fe94bf5e --- /dev/null +++ b/examples/radixshmem_configs/radix_server_single_node.sh @@ -0,0 +1,11 @@ +#!/bin/bash +# The operator's radix-server for one node, no cluster: nothing about the model +# on the command line. FlexKV's first client brings the geometry (page size, +# bytes per block, SWA page + window); the server plans the slot counts from the +# budget and FlexKV adopts them (docs/radixshmem/config_zh.md section 3). +# --swa-ratio is needed only for models with an SWA pool (DeepSeek-V4 etc.). +set -eu +radix-server --name "${NAME:-/flexkv}" \ + --data-bytes "${DATA_BYTES:-64G}" \ + ${SWA_RATIO:+--swa-ratio "$SWA_RATIO"} \ + ${HUGEPAGE_PATH:+--hugepage-path "$HUGEPAGE_PATH"} diff --git a/examples/radixshmem_configs/radixshmem.yaml b/examples/radixshmem_configs/radixshmem.yaml new file mode 100644 index 000000000..6a8279514 --- /dev/null +++ b/examples/radixshmem_configs/radixshmem.yaml @@ -0,0 +1,19 @@ +# FlexKV side of radixshmem mode: which radix-server to attach to and the +# prefetch limits. Every key is at its default; the file is equivalent to not +# setting FLEXKV_RADIXSHMEM_CONFIG_PATH at all. +# +# export FLEXKV_ENABLE_RADIXSHMEM=1 +# export FLEXKV_CPU_LAYOUT=BLOCKFIRST +# export FLEXKV_RADIXSHMEM_CONFIG_PATH=$PWD/examples/radixshmem_configs/radixshmem.yaml +# +# The server itself is the operator's process (radix_server_single_node.sh / +# radix_server_multi_node.sh). FlexKV hands it the geometry and adopts the slot +# counts it plans from --data-bytes. Reference: docs/radixshmem/config_zh.md +server: + name: /flexkv # radix-server --name + endpoint: "" # "" = unix:///dev/shm/flexkv.sock + ready_timeout_s: 600 # reachable AND ready (prefault, cluster rendezvous) +client: + prefetch_timeout_ms: 5000 + prefetch_max_inflight: 128 + max_outstanding: 256 diff --git a/examples/radixshmem_configs/radixshmem_multi_node.yaml b/examples/radixshmem_configs/radixshmem_multi_node.yaml deleted file mode 100644 index ca19e731a..000000000 --- a/examples/radixshmem_configs/radixshmem_multi_node.yaml +++ /dev/null @@ -1,24 +0,0 @@ -# FlexKV radixshmem mode, a 4-node cluster sharing CPU KV over RDMA. -# -# export FLEXKV_ENABLE_RADIXSHMEM=1 -# export FLEXKV_RADIXSHMEM_CONFIG_PATH=/etc/flexkv/radixshmem.yaml -# export FLEXKV_CPU_LAYOUT=BLOCKFIRST -# -# This file is GLOBAL: copy it byte for byte to every node. Each node resolves -# its own IP from cluster.rpc_interface and gets its identity and rank from the -# etcd rendezvous. Only the keys a cluster must set are listed; everything else -# keeps its default. Reference: docs/radixshmem/config_zh.md - -cluster: - cluster_id: prod_a - expected_min_nodes: 4 - num_rht_shards: 4 - registry: etcd://10.0.0.1:2379 - rpc_interface: bond0 - index_dev: mlx5_bond_0 - -data: - transfer_devices: [mlx5_0, mlx5_1, mlx5_2, mlx5_3, mlx5_4, mlx5_5, mlx5_6, mlx5_7] - -index: - background_evict_ratio: 0.05 diff --git a/examples/radixshmem_configs/radixshmem_single_node.yaml b/examples/radixshmem_configs/radixshmem_single_node.yaml deleted file mode 100644 index 2bc9ba649..000000000 --- a/examples/radixshmem_configs/radixshmem_single_node.yaml +++ /dev/null @@ -1,17 +0,0 @@ -# FlexKV radixshmem mode, one node (no etcd, no RDMA). -# -# export FLEXKV_ENABLE_RADIXSHMEM=1 -# export FLEXKV_RADIXSHMEM_CONFIG_PATH=$PWD/examples/radixshmem_configs/radixshmem_single_node.yaml -# export FLEXKV_CPU_LAYOUT=BLOCKFIRST -# -# Every key below is at its default; the file is equivalent to setting no -# FLEXKV_RADIXSHMEM_CONFIG_PATH at all. It is here to show what is tunable on a -# single node. Slot counts / slot bytes / shm names are derived from the FlexKV -# cache configuration and must not appear here. Reference: docs/radixshmem/config_zh.md - -cluster: - cluster_id: flexkv - expected_min_nodes: 0 - -index: - background_evict_ratio: 0.05 diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 79a83c61d..1a0692142 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -2,9 +2,9 @@ # cython: boundscheck=True, wraparound=True """ The radixshmem CPU tier: one process's `shmradix.RadixClient` on the node's -radix-server (index shm + SlotStore + RDMA engine, see -`flexkv.server.shm_radix_bootstrap`). Used only by -`flexkv.cache.radix_shmem_planner.RadixShmemCacheEngine`. +radix-server (the operator's process: index shm + SlotStore + RDMA engine; the +attach and the geometry hand-off are in `flexkv.server.shm_radix_bootstrap`). +Used only by `flexkv.cache.radix_shmem_planner.RadixShmemCacheEngine`. Contracts: @@ -147,18 +147,23 @@ class CacheEngineRadixShmem: (one per DP scheduler process) share a server and operate concurrently.""" def __init__(self, - shm_name: str, + server_name: str, *, tokens_per_block: int, num_total_blocks: int, - peer_enabled: bool = False, + geometry: Any = None, + peer_enabled: Optional[bool] = None, swa_config: Optional[SWAPoolConfig] = None, event_collector: Optional[KVEventCollector] = None, metrics_collector=None): - """`shm_name` is the base index name (`radix_index_name`); the server - must be running or starting. `peer_enabled` only takes effect on a - clustered region. `num_total_blocks` is FlexKV's expectation; the - region's capacity is authoritative.""" + """`server_name` is the radix-server's ``--name``; the server is the + operator's process, running but not necessarily ready. With + `geometry` (FlexKV's `RadixGeometry` or a `shmradix.Geometry`) the + attach hands it FlexKV's slot shape (idempotent) and waits for it to + come up. `peer_enabled` None = follow the region: peer reuse whenever + the server is part of a cluster; False switches it off. + `num_total_blocks` is FlexKV's expectation; the region's capacity is + authoritative.""" from flexkv.server.shm_radix_bootstrap import attach_radix_client self.event_collector = event_collector @@ -166,10 +171,13 @@ def __init__(self, cpu_swa = swa_config.for_cache_tier(DeviceType.CPU) if swa_config is not None else None self.swa_enabled = cpu_swa is not None and cpu_swa.num_slots > 0 - self._client = attach_radix_client(shm_name) + self._client = attach_radix_client(server_name, geometry=geometry, + label="CacheEngineRadixShmem") self._tree = self._client # index ops pass through the client self.shm_name = self._client.info.index_name # node-suffixed when distributed self.is_distributed = bool(self._client.is_distributed()) + if peer_enabled is None: + peer_enabled = self.is_distributed self.peer_enabled = bool(peer_enabled) and self.is_distributed if peer_enabled and not self.is_distributed: flexkv_logger.warning( diff --git a/flexkv/cache/radix_shmem_planner.py b/flexkv/cache/radix_shmem_planner.py index eaa129d83..ebcfd5bce 100644 --- a/flexkv/cache/radix_shmem_planner.py +++ b/flexkv/cache/radix_shmem_planner.py @@ -20,7 +20,7 @@ Not supported here: the SSD and REMOTE tiers, the Redis-backed P2P paths (``enable_p2p_cpu`` / ``enable_p2p_ssd``) and kv sharing. Peer reuse follows -the radixshmem YAML instead (`RadixShmemConfig.distributed`). +the radix-server instead (on whenever it is part of a cluster). """ from __future__ import annotations @@ -146,8 +146,8 @@ def _check_cache_config(cache_config: CacheConfig) -> None: if cache_config.enable_p2p_cpu or cache_config.enable_p2p_ssd or cache_config.enable_kv_sharing: raise ValueError( "radix_shmem does its own peer reuse (etcd + RDMA inside the " - "radix-server, on whenever the radixshmem YAML makes the cluster " - "distributed); enable_p2p_cpu / enable_p2p_ssd must be off") + "radix-server, on whenever it is started with cluster flags); " + "enable_p2p_cpu / enable_p2p_ssd must be off") class RadixShmemCacheEngine(GlobalCacheEngine): @@ -173,18 +173,20 @@ def _build_cpu_cache_engine(self, event_collector: Optional[KVEventCollector]): """Attach to this node's radix-server as a `RadixClient`. - The server (index + SlotStore + peer transfer) is brought up by the - KVManager bootstrap process or by the operator (`shm_radix_bootstrap`); - the attach waits for it to be ready. + The server (index + SlotStore + peer transfer) is the operator's + `radix-server` process. The attach brings FlexKV's geometry + (idempotent: the KVManager already handed it over and adopted the slot + counts into `cache_config`) and waits for the server to be ready. Peer + reuse follows the server: on whenever it is part of a cluster. """ - from flexkv.server.shm_radix_bootstrap import radix_index_name + from flexkv.server.shm_radix_bootstrap import expected_geometry rcfg = self._radix_config return CacheEngineRadixShmem( - radix_index_name(rcfg.local_id), + rcfg.server_name, + geometry=expected_geometry(self.model_config, cache_config), tokens_per_block=cache_config.tokens_per_block, num_total_blocks=cache_config.num_cpu_blocks, - peer_enabled=rcfg.distributed, swa_config=cache_config.swa, event_collector=event_collector, metrics_collector=self._metrics_collector, diff --git a/flexkv/common/config.py b/flexkv/common/config.py index 2ddc307c7..ae9183e12 100644 --- a/flexkv/common/config.py +++ b/flexkv/common/config.py @@ -840,20 +840,13 @@ def __str__(self) -> str: server_recv_port=os.getenv('FLEXKV_SERVER_RECV_PORT', 'ipc:///tmp/flexkv_server'), # radixshmem mode: the CPU tier is radixshmem's index + SlotStore, one - # radix-server per node, one shared TE, a KVTaskEngine per DP process (no - # KVServer). Everything else about that mode -- cluster membership, RDMA - # devices, prefetch limits -- is the YAML at FLEXKV_RADIXSHMEM_CONFIG_PATH + # radix-server per node (a process the operator starts: `radix-server + # --name /flexkv --data-bytes ...`), one shared TE, a KVTaskEngine per DP + # process (no KVServer). FlexKV only attaches to the server; which one and + # the prefetch limits are the YAML at FLEXKV_RADIXSHMEM_CONFIG_PATH # (flexkv.common.radixshmem_config; reference docs/radixshmem/config_zh.md). enable_radixshmem=bool(int(os.getenv('FLEXKV_ENABLE_RADIXSHMEM', 0))), radixshmem_config_path=os.getenv('FLEXKV_RADIXSHMEM_CONFIG_PATH', '') or None, - # embedded: the bootstrap DP process launches the radix-server subprocess; - # external: a radix-server started by the operator is attached to. - radix_server_launch_mode=os.getenv('FLEXKV_RADIX_SERVER_LAUNCH_MODE', 'embedded').lower(), - # Per-node overrides of the global YAML, for several nodes on one host - # (tests): the node's etcd identity and the bootstrap IP peers dial. Unset - # in a real deployment, where both derive from cluster.rpc_interface. - radix_node_name=os.getenv('FLEXKV_RADIX_NODE_NAME', ''), - radix_rpc_address=os.getenv('FLEXKV_RADIX_RPC_ADDRESS', ''), index_accel=bool(int(os.getenv('FLEXKV_INDEX_ACCEL', 1))), cpu_layout_type=KVCacheLayoutType(os.getenv('FLEXKV_CPU_LAYOUT', 'BLOCKFIRST').upper()), diff --git a/flexkv/common/radixshmem_config.py b/flexkv/common/radixshmem_config.py index 900b57edf..9991f1586 100644 --- a/flexkv/common/radixshmem_config.py +++ b/flexkv/common/radixshmem_config.py @@ -2,26 +2,31 @@ # cython: boundscheck=True, wraparound=True """The radixshmem-mode configuration file (``FLEXKV_RADIXSHMEM_CONFIG_PATH``). -One YAML, identical on every node of a cluster, with five sections: - - cluster / data / index / server - Passed through by key to ``shmradix.ClusterConfig`` / - ``DataPlaneConfig`` / ``IndexConfig`` / ``RadixServerConfig``. Keys are - validated against the dataclass fields of the installed shmradix, so a - new radixshmem field is configurable without a FlexKV change and a typo - fails at startup. Geometry fields (slot counts, slot bytes, alignment, - shm names) are derived from ``CacheConfig`` and rejected here. +In radixshmem mode the CPU tier is a ``radix-server`` process the operator +starts on every node (``radix-server --name /flexkv --data-bytes 64G ...``): +the index shm, the SlotStore, the transfer engine and the cluster membership +all belong to that process and are set on its command line. FlexKV never +creates a server; it attaches a ``shmradix.RadixClient``. So this file holds +the two things FlexKV has to know, and nothing else: + + server + Which radix-server to attach to: its ``--name`` (which also derives the + default gRPC socket ``unix:///dev/shm/.sock`` and the prefix of + FlexKV's own TE channels), an ``endpoint`` override, and how long a + FlexKV process waits for the server to exist and become ready. client - FlexKV's own RadixClient / prefetch settings. + FlexKV's RadixClient / prefetch settings. -Per-node values do not belong in a global file: ``cluster.node_name`` and -``cluster.rpc_address`` are rejected. A node derives its identity from the IP -``cluster.rpc_interface`` resolves to; the two environment variables -``FLEXKV_RADIX_NODE_NAME`` / ``FLEXKV_RADIX_RPC_ADDRESS`` override that for -several nodes on one host (tests). +The slot geometry (tokens per block, bytes of one CPU block and one SWA page, +the SWA window) is derived from ``ModelConfig`` / ``CacheConfig`` and handed to +the server by FlexKV's clients (``flexkv.server.shm_radix_bootstrap``); the +slot COUNTS come back from the server, which plans them from its byte budget. +None of that is in this file. A file that still carries the former ``cluster`` +/ ``data`` / ``index`` sections is rejected with a pointer to the +``radix-server`` flags they moved to. -``cluster.cluster_id`` is the only namespace: the etcd key prefix and, through -:meth:`RadixShmemConfig.local_id`, every shm / socket / IPC name on this host. +No YAML at all is a valid configuration: it attaches to ``radix-server --name +/flexkv`` on the local socket. Reference: ``docs/radixshmem/config_zh.md``. """ @@ -29,54 +34,38 @@ import dataclasses import threading -from typing import Any, Dict, Optional, Set, Tuple +from typing import Any, Dict, Optional, Tuple import yaml from flexkv.common.config import GLOBAL_CONFIG_FROM_ENV -SECTIONS = ("cluster", "data", "index", "server", "client") - -# Values FlexKV sets differently from radixshmem's own defaults; anything not -# listed takes the shmradix dataclass default. One more default depends on the -# geometry and is resolved in shm_radix_bootstrap.build_radix_server_config: -# index.register_chunk_size = REGISTER_CHUNK_TOKENS // tokens_per_block, so an -# RHT registration chunk covers REGISTER_CHUNK_TOKENS tokens whatever the block -# size (radixshmem's own default is 128 blocks). -REGISTER_CHUNK_TOKENS = 4096 - -FLEXKV_DEFAULTS: Dict[str, Dict[str, Any]] = { - "cluster": { - "cluster_id": "flexkv", - "bootstrap_timeout_sec": 120, - # 1 is a blind overwrite that loses routing entries. - "rht_slots_per_bucket": 4, - }, - "data": {}, - "index": {"data_pool_ratio": 8.0}, - "server": {}, -} - -# Derived from CacheConfig / ModelConfig (shm_radix_bootstrap.expected_geometry) -# or per node; rejected in the file. -FORBIDDEN_KEYS: Dict[str, Set[str]] = { - "cluster": {"node_name", "rpc_address"}, - "data": {"data_bytes", "full_slot_bytes", "swa_slot_bytes", "mamba_slot_bytes", - "slot_align", "data_name"}, - "index": {"name", "tokens_per_block", "full_slots", "swa_slots", "swa_window_blocks", - "mamba_slots", "evict_policy"}, - "server": set(), -} - -_RDMA_TRANSPORTS = {"xrc", "dc"} -_REMOTE_OP_TRANSPORTS = {"zmq", "dc"} -_RHT_SLOTS = {1, 2, 4, 8} +SECTIONS = ("server", "client") +# Sections of the previous file format. Their keys are radix-server flags now. +RETIRED_SECTIONS = ("cluster", "data", "index") + +DEFAULT_SERVER_NAME = "/flexkv" class RadixShmemConfigError(ValueError): """The file is not a valid radixshmem-mode configuration.""" +@dataclasses.dataclass(frozen=True) +class RadixServerSettings: + """Which radix-server this node's FlexKV attaches to.""" + # ``radix-server --name``: the index shm name. Also the default socket + # (``unix:///dev/shm/.sock``) and, sanitized, the prefix of FlexKV's + # TE shm channels on this host (``RadixShmemConfig.te_server_id``). + name: str = DEFAULT_SERVER_NAME + # gRPC endpoint; "" = the default socket derived from ``name``. + endpoint: str = "" + # How long a FlexKV process waits for the server to be reachable AND ready. + # Covers the operator starting it late, the SlotStore prefault and, on a + # cluster, the rendezvous (the server's --bootstrap-timeout). + ready_timeout_s: float = 600.0 + + @dataclasses.dataclass(frozen=True) class RadixClientSettings: """FlexKV-side settings of the RadixClient and the prefetch path.""" @@ -93,108 +82,47 @@ class RadixClientSettings: @dataclasses.dataclass(frozen=True) class RadixShmemConfig: path: Optional[str] - cluster: Dict[str, Any] - data: Dict[str, Any] - index: Dict[str, Any] - server: Dict[str, Any] + server: RadixServerSettings = RadixServerSettings() client: RadixClientSettings = RadixClientSettings() - # ----------------------------------------------------------- cluster - @property - def cluster_id(self) -> str: - return str(self.cluster["cluster_id"]) - - @property - def node_name(self) -> str: - return str(self.cluster.get("node_name", "")) - - @property - def rpc_address(self) -> str: - return str(self.cluster.get("rpc_address", "")) - - @property - def expected_min_nodes(self) -> int: - return int(self.cluster.get("expected_min_nodes", 0)) - - @property - def num_rht_shards(self) -> int: - return int(self.cluster.get("num_rht_shards", 0)) - - @property - def distributed(self) -> bool: - """radixshmem's own criterion (ClusterConfig.distributed).""" - return self.expected_min_nodes > 1 or self.num_rht_shards > 1 - - @property - def bootstrap_timeout_sec(self) -> float: - return float(self.cluster.get("bootstrap_timeout_sec", 60)) - + # ------------------------------------------------------------ server @property - def attach_timeout_s(self) -> float: - """How long a FlexKV process waits for the radix-server: the cluster - rendezvous plus a margin for SlotStore creation / prefault.""" - return self.bootstrap_timeout_sec + 60.0 + def server_name(self) -> str: + return self.server.name @property - def local_id(self) -> str: - """Prefix of every name on this host: the index / SlotStore shm, the - gRPC socket, FlexKV's TE channels and GPU registration port. The - cluster id, suffixed with the node name when one was given so that - co-located nodes do not share regions.""" - return f"{self.cluster_id}_{self.node_name}" if self.node_name else self.cluster_id + def endpoint(self) -> str: + """gRPC endpoint; "" = radixshmem's ``unix:///dev/shm/.sock``.""" + return self.server.endpoint - # ------------------------------------------------------------ server @property - def endpoint(self) -> str: - """gRPC endpoint; "" = radixshmem's unix:///dev/shm/.sock.""" - return str(self.server.get("endpoint", "")) + def ready_timeout_s(self) -> float: + return float(self.server.ready_timeout_s) @property - def hugepage_path(self) -> str: - return str(self.server.get("hugepage_path", "")) + def te_server_id(self) -> str: + """Prefix of FlexKV's own IPC objects on this host (the TE control + block and channels): the server name without its leading slash, so + two FlexKV deployments on one host that attach to different servers + never share a channel.""" + return self.server.name.lstrip("/").replace("/", "_") # ------------------------------------------------------------- tests - def replace_cluster(self, **changes: Any) -> "RadixShmemConfig": - """A copy with ``cluster`` keys changed (test helper; bypasses the - forbidden-key check so node_name / rpc_address can be set).""" - return dataclasses.replace(self, cluster={**self.cluster, **changes}) - def replace_server(self, **changes: Any) -> "RadixShmemConfig": - return dataclasses.replace(self, server={**self.server, **changes}) + return dataclasses.replace(self, server=dataclasses.replace(self.server, **changes)) + + def replace_client(self, **changes: Any) -> "RadixShmemConfig": + return dataclasses.replace(self, client=dataclasses.replace(self.client, **changes)) def describe(self) -> str: where = self.path or "(defaults)" - s = f"{where}: cluster_id={self.cluster_id}" - if self.distributed: - s += (f", expected_min_nodes={self.expected_min_nodes}, " - f"registry={self.cluster.get('registry')}, " - f"rpc_interface={self.cluster.get('rpc_interface') or '-'}, " - f"rpc_address={self.rpc_address or '-'}, node_name={self.node_name or '(auto)'}") - return s + return (f"{where}: radix-server {self.server_name} " + f"(endpoint={self.endpoint or 'unix:///dev/shm/' + self.te_server_id + '.sock'}, " + f"ready_timeout_s={self.ready_timeout_s:.0f})") # ------------------------------------------------------------------ loading -def _shmradix_dataclasses(): - try: - import shmradix - except ImportError as exc: # pragma: no cover - raise ImportError( - "shmradix is not installed; install it from the radixshmem repo " - "(pip install -e radixshmem/python)") from exc - try: - return { - "cluster": shmradix.ClusterConfig, - "data": shmradix.DataPlaneConfig, - "index": shmradix.IndexConfig, - "server": shmradix.RadixServerConfig, - } - except AttributeError as exc: - raise ImportError( - "shmradix lacks the RadixServer configuration dataclasses: FlexKV needs " - "the RadixServer / RadixClient surface of radixshmem") from exc - - def _read_yaml(path: str) -> Dict[str, Any]: with open(path) as f: loaded = yaml.safe_load(f) @@ -214,86 +142,36 @@ def _section(raw: Dict[str, Any], name: str, path: str) -> Dict[str, Any]: return dict(sec) -def _as_list(value: Any) -> Any: - if isinstance(value, str): - return [v.strip() for v in value.split(",") if v.strip()] - return value - - -def _passthrough_section(name: str, given: Dict[str, Any], dc, path: str) -> Dict[str, Any]: - """FlexKV defaults overlaid with the file's keys, validated against the - shmradix dataclass ``dc``.""" - fields = {f.name for f in dataclasses.fields(dc)} - if name == "server": - # RadixServerConfig's nested sections are configured by their own - # sections here, not inline. - fields -= {"index", "data", "cluster"} - forbidden = FORBIDDEN_KEYS[name] & set(given) - if forbidden: - raise RadixShmemConfigError( - f"{path}: '{name}.{sorted(forbidden)[0]}' is not configurable: " - + ("geometry is derived from the FlexKV cache configuration" - if name in ("data", "index") else - "it is a per-node value; set FLEXKV_RADIX_NODE_NAME / " - "FLEXKV_RADIX_RPC_ADDRESS on that node instead")) - unknown = set(given) - fields - if unknown: - raise RadixShmemConfigError( - f"{path}: unknown key(s) in '{name}': {sorted(unknown)}; " - f"shmradix.{dc.__name__} has {sorted(fields)}") - merged = {**FLEXKV_DEFAULTS[name], **given} - if "transfer_devices" in merged: - merged["transfer_devices"] = [str(d) for d in _as_list(merged["transfer_devices"])] - if "rht_shard_holders" in merged: - merged["rht_shard_holders"] = [int(r) for r in _as_list(merged["rht_shard_holders"])] - return merged - - -def _client_section(given: Dict[str, Any], path: str) -> RadixClientSettings: - fields = {f.name for f in dataclasses.fields(RadixClientSettings)} - unknown = set(given) - fields +def _typed_section(name: str, given: Dict[str, Any], dc, path: str): + """``dc(**given)`` after checking the keys and coercing the value types.""" + fields = {f.name: f for f in dataclasses.fields(dc)} + unknown = set(given) - set(fields) if unknown: raise RadixShmemConfigError( - f"{path}: unknown key(s) in 'client': {sorted(unknown)}; expected {sorted(fields)}") - return RadixClientSettings(**{k: int(v) for k, v in given.items()}) + f"{path}: unknown key(s) in '{name}': {sorted(unknown)}; expected {sorted(fields)}") + values: Dict[str, Any] = {} + for key, value in given.items(): + typ = fields[key].type + try: + if typ in ("int", int): + values[key] = int(value) + elif typ in ("float", float): + values[key] = float(value) + else: + values[key] = "" if value is None else str(value) + except (TypeError, ValueError) as exc: + raise RadixShmemConfigError(f"{path}: '{name}.{key}' has an invalid value {value!r}") from exc + return dc(**values) def _validate(cfg: RadixShmemConfig, path: str) -> None: - c = cfg.cluster - if not cfg.cluster_id: - raise RadixShmemConfigError(f"{path}: cluster.cluster_id must not be empty") - if cfg.distributed: - if not c.get("registry"): - raise RadixShmemConfigError( - f"{path}: cluster mode (expected_min_nodes > 1) needs cluster.registry, " - f"e.g. 'etcd://10.0.0.1:2379'") - if not c.get("rpc_interface") and not cfg.rpc_address: - raise RadixShmemConfigError( - f"{path}: cluster mode needs cluster.rpc_interface (the NIC whose IP peers " - f"dial and this node's identity derives from) or FLEXKV_RADIX_RPC_ADDRESS") - if cfg.rpc_address == "0.0.0.0": - raise RadixShmemConfigError( - "FLEXKV_RADIX_RPC_ADDRESS=0.0.0.0 gives every node the same identity; " - "use this node's address") - if cfg.expected_min_nodes > 0 and cfg.num_rht_shards > cfg.expected_min_nodes: - raise RadixShmemConfigError( - f"{path}: cluster.num_rht_shards={cfg.num_rht_shards} exceeds " - f"expected_min_nodes={cfg.expected_min_nodes}; there cannot be more RHT shard " - f"holders than nodes") - slots = int(c.get("rht_slots_per_bucket", 1)) - if slots not in _RHT_SLOTS: + name = cfg.server_name + if not name or not name.startswith("/") or len(name) < 2 or any(c.isspace() for c in name): raise RadixShmemConfigError( - f"{path}: cluster.rht_slots_per_bucket={slots} must be one of {sorted(_RHT_SLOTS)}") - for key in ("rht_transport", "peer_index_transport"): - val = c.get(key) - if val is not None and val not in _RDMA_TRANSPORTS: - raise RadixShmemConfigError( - f"{path}: cluster.{key}={val!r} must be one of {sorted(_RDMA_TRANSPORTS)}") - rot = c.get("remote_op_transport") - if rot is not None and rot not in _REMOTE_OP_TRANSPORTS: - raise RadixShmemConfigError( - f"{path}: cluster.remote_op_transport={rot!r} must be one of " - f"{sorted(_REMOTE_OP_TRANSPORTS)}") + f"{path}: server.name={name!r} must be a shm name that starts with '/' " + f"(the radix-server's --name, e.g. '/flexkv')") + if cfg.ready_timeout_s <= 0: + raise RadixShmemConfigError(f"{path}: server.ready_timeout_s must be > 0") if cfg.client.prefetch_max_inflight >= cfg.client.max_outstanding: raise RadixShmemConfigError( f"{path}: client.prefetch_max_inflight={cfg.client.prefetch_max_inflight} must be " @@ -302,31 +180,29 @@ def _validate(cfg: RadixShmemConfig, path: str) -> None: raise RadixShmemConfigError(f"{path}: client.prefetch_timeout_ms must be > 0") -def load_radixshmem_config(path: Optional[str] = None, - *, - node_name: str = "", - rpc_address: str = "") -> RadixShmemConfig: - """Parse ``path`` (None or "" = all defaults, i.e. standalone) and apply - the per-node overrides. Raises :class:`RadixShmemConfigError` on an - invalid file, ``ImportError`` without shmradix.""" - dcs = _shmradix_dataclasses() +def load_radixshmem_config(path: Optional[str] = None) -> RadixShmemConfig: + """Parse ``path`` (None or "" = all defaults: ``radix-server --name /flexkv`` + on the local socket). Raises :class:`RadixShmemConfigError` on an invalid + file.""" label = path or "(defaults)" raw = _read_yaml(path) if path else {} + retired = [s for s in RETIRED_SECTIONS if s in raw] + if retired: + raise RadixShmemConfigError( + f"{label}: section(s) {retired} are not FlexKV's any more: the radix-server owns " + f"its cluster, data plane and index settings and takes them on its command line " + f"(radix-server --data-bytes / --swa-ratio / --expected-min-nodes / --registry / " + f"--transfer-dev ...). FlexKV only attaches to it; keep 'server' and 'client' here. " + f"See docs/radixshmem/config_zh.md") unknown = set(raw) - set(SECTIONS) if unknown: raise RadixShmemConfigError( f"{label}: unknown section(s) {sorted(unknown)}; expected {list(SECTIONS)}") - sections = {name: _passthrough_section(name, _section(raw, name, label), dcs[name], label) - for name in ("cluster", "data", "index", "server")} - if node_name: - sections["cluster"]["node_name"] = str(node_name) - if rpc_address: - sections["cluster"]["rpc_address"] = str(rpc_address) - # radixshmem lets the interface win over the address; an explicit - # per-node address means the global interface must not apply here. - sections["cluster"]["rpc_interface"] = "" - cfg = RadixShmemConfig(path=path or None, client=_client_section( - _section(raw, "client", label), label), **sections) + cfg = RadixShmemConfig( + path=path or None, + server=_typed_section("server", _section(raw, "server", label), RadixServerSettings, label), + client=_typed_section("client", _section(raw, "client", label), RadixClientSettings, label), + ) _validate(cfg, label) return cfg @@ -334,26 +210,22 @@ def load_radixshmem_config(path: Optional[str] = None, # -------------------------------------------------------------- singleton _lock = threading.Lock() -_cached: Optional[Tuple[Tuple[str, str, str], RadixShmemConfig]] = None +_cached: Optional[Tuple[Tuple[str], RadixShmemConfig]] = None -def _env_key() -> Tuple[str, str, str]: - env = GLOBAL_CONFIG_FROM_ENV - return (str(env.radixshmem_config_path or ""), str(env.radix_node_name or ""), - str(env.radix_rpc_address or "")) +def _env_key() -> Tuple[str]: + return (str(GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path or ""),) def get_radixshmem_config() -> RadixShmemConfig: """The process's configuration: loaded from ``GLOBAL_CONFIG_FROM_ENV`` - (``FLEXKV_RADIXSHMEM_CONFIG_PATH`` + the two per-node overrides) on first - use and whenever those three values change.""" + (``FLEXKV_RADIXSHMEM_CONFIG_PATH``) on first use and whenever that value + changes.""" global _cached key = _env_key() with _lock: if _cached is None or _cached[0] != key: - path, node_name, rpc_address = key - _cached = (key, load_radixshmem_config(path or None, node_name=node_name, - rpc_address=rpc_address)) + _cached = (key, load_radixshmem_config(key[0] or None)) return _cached[1] diff --git a/flexkv/integration/sglang/connector.py b/flexkv/integration/sglang/connector.py index 75031c131..18b938b5b 100644 --- a/flexkv/integration/sglang/connector.py +++ b/flexkv/integration/sglang/connector.py @@ -63,9 +63,12 @@ def _radixshmem_distributed() -> bool: - """Whether the radixshmem YAML describes a cluster (peer pulls possible).""" - from flexkv.common.radixshmem_config import get_radixshmem_config - return get_radixshmem_config().distributed + """Whether this node's radix-server is part of a cluster (peer pulls + possible). The server is the operator's process and knows; every TP rank + asks it the same question once it is ready, so the prefetch gate below is + the same in all ranks (the PREFETCH_START scatter needs that).""" + from flexkv.server.shm_radix_bootstrap import radix_server_is_distributed + return radix_server_is_distributed(label="FlexKVConnector") logger = logging.getLogger(__name__) diff --git a/flexkv/kvmanager.py b/flexkv/kvmanager.py index baf5496b0..77444d909 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -78,20 +78,19 @@ def __init__(self, ) if self.enable_radixshmem and (cache_config.enable_p2p_cpu or cache_config.enable_p2p_ssd): - # Peer reuse is the radix-server's (etcd + RDMA), switched on by the - # radixshmem YAML making the cluster distributed; the Redis-backed - # P2P paths these flags select must stay off. + # Peer reuse is the radix-server's (etcd + RDMA), switched on by + # its cluster flags (--expected-min-nodes / --registry); the + # Redis-backed P2P paths these flags select must stay off. raise ValueError( "radix_shmem does its own peer reuse; set enable_p2p_cpu=False " "and enable_p2p_ssd=False (cross-node reuse follows the " - "radixshmem YAML: expected_min_nodes / num_rht_shards)" + "radix-server's cluster flags)" ) - # Prefix of this host's radix regions and TE channels: the YAML's - # cluster_id (plus the node name when several nodes share the host). + # Prefix of this host's TE channels: the attached radix-server's name. self._shm_radix_id = None if self.enable_radixshmem: from flexkv.common.radixshmem_config import get_radixshmem_config - self._shm_radix_id = get_radixshmem_config().local_id + self._shm_radix_id = get_radixshmem_config().te_server_id flexkv_logger.info( f"[KVManager] IPC ports: server_recv_port={self.server_recv_port}, " @@ -139,10 +138,8 @@ def __init__(self, self.redis_meta_client = None self.enable_mps = GLOBAL_CONFIG_FROM_ENV.enable_mps self.owns_mps = self.enable_mps and self.server_launch_mode != "external" - # The embedded radix-server subprocess — only the bootstrap process - # holds this; others have None. - self._shm_radix_server = None - # TE-process handle — only the bootstrap process holds this. + # TE-process handle — only the bootstrap process holds this. The + # radix-server itself is the operator's process, not FlexKV's. self._shm_te_process = None # Local KVTaskEngine for the radix-shmem path (per-DP). self.kv_task_engine = None @@ -198,26 +195,46 @@ def _init_radix_shmem_path(self, event_collector: Optional[KVEventCollector]) -> None: """Initialize the radix-shmem multi-DP path. - Everything shared by this inference instance's DP processes on the - node — the radix shm regions and the single TE subprocess — is set up - by the node-local bootstrap proc (local DP client 0) only. Every other - proc builds its own KVTaskEngine and - attaches: `CacheEngineRadixShmem` polls for its region, and the TE - channel handle blocks in `ShmControlBlock.wait_ready`. + The node's radix-server is a process the operator started + (``radix-server --name --data-bytes ...``); FlexKV never + creates one. Every DP process attaches to it with FlexKV's geometry + (the first one configures the server, the others find it configured), + takes over the slot counts the server planned from its byte budget, + and builds its own KVTaskEngine on it. The single TE subprocess all DP + processes of this node share is spawned by the node-local bootstrap + proc (local DP client 0) only; every other proc attaches to its channel + (``ShmControlBlock.wait_ready``). Each CE process gets a disjoint graph/op id range so submissions to the single shared TE never collide. """ from flexkv.common.transfer import TransferOp + from flexkv.server.shm_radix_bootstrap import (adopt_geometry, attach_radix_client, + expected_geometry, radix_cluster_rank) TransferOpGraph.set_graph_id_range(self.dp_client_id << 32, (self.dp_client_id + 1) << 32) TransferOp.set_op_id_range(self.dp_client_id << 32, (self.dp_client_id + 1) << 32) + # Hand the server FlexKV's geometry and take its slot counts BEFORE + # anything sizes a pool from cache_config: the TE's StorageEngine and + # the cache engine both read num_cpu_blocks / swa.num_slots. + geometry = expected_geometry(self.model_config, self.cache_config) + client = attach_radix_client(geometry=geometry, label="KVManager") + try: + adopt_geometry(self.cache_config, client, label="KVManager") + self.cache_config.distributed_node_id = radix_cluster_rank(client) + flexkv_logger.info( + f"[kv manager] radix-server {client.name}: cluster rank " + f"{self.cache_config.distributed_node_id}/{client.info.world_size}, " + f"FlexKV geometry {geometry.describe()}") + finally: + client.close() + try: if self.local_dp_client_id == 0: - self._bootstrap_radix_shmem() + self._spawn_shm_te() # KVTaskEngine reads GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and builds a # RadixShmemCacheEngine (CPU tier = RadixClient on the radix-server). @@ -230,9 +247,8 @@ def _init_radix_shmem_path(self, shm_te_channel_id=self.local_dp_client_id, ) except BaseException: - # A failure after the TE / radix-server subprocesses were spawned - # must not leave them running (the TE would wait for GPU - # registrations forever). + # A failure after the TE subprocess was spawned must not leave it + # running (it would wait for GPU registrations forever). self._shutdown_radix_shmem_children() raise @@ -240,45 +256,13 @@ def _shutdown_radix_shmem_children(self) -> None: if self._shm_te_process is not None: self._shm_te_process.shutdown() self._shm_te_process = None - if self._shm_radix_server is not None: - self._shm_radix_server.shutdown() - self._shm_radix_server = None - - def _bootstrap_radix_shmem(self) -> None: - """Bootstrap proc (dp 0) only: bring up this node's radix-server (index + - SlotStore + peer transfer) and spawn the shared TE. - - The server is a subprocess (``FLEXKV_RADIX_SERVER_LAUNCH_MODE=embedded``) - or one the operator started (``external``); either way every FlexKV - process attaches by name. Peer reuse needs no Redis address book any - more: the server resolves peers through etcd and pulls their blocks - itself (``RadixClient.pull_async`` from the prefetch path).""" - from flexkv.server.shm_radix_bootstrap import (RadixServerProcess, - build_radix_server_config, - radix_socket_path) - from flexkv.transfer_manager import TransferManagerShmTEProcess - launch_mode = GLOBAL_CONFIG_FROM_ENV.radix_server_launch_mode - if launch_mode not in ("embedded", "external"): - raise ValueError( - "FLEXKV_RADIX_SERVER_LAUNCH_MODE must be embedded or external, " - f"got {launch_mode!r}" - ) - if launch_mode == "embedded": - server_cfg = build_radix_server_config(self.model_config, self.cache_config) - self._shm_radix_server = RadixServerProcess(server_cfg).start() - self.cache_config.distributed_node_id = int( - self._shm_radix_server.cluster_rank) - flexkv_logger.info( - f"[kv manager] radix-server for {self._shm_radix_id} is up: " - f"cluster rank {self.cache_config.distributed_node_id}" - ) - else: - from flexkv.common.radixshmem_config import get_radixshmem_config - flexkv_logger.info( - f"[kv manager] attaching to an external radix-server at " - f"{get_radixshmem_config().endpoint or radix_socket_path(self._shm_radix_id)}" - ) + def _spawn_shm_te(self) -> None: + """Bootstrap proc (local dp 0) only: spawn the TE subprocess every DP + process of this node feeds over its shm channel. The TE attaches to + the same radix-server (its SlotStore is the CPU pool) with the cache + config whose slot counts were just adopted.""" + from flexkv.transfer_manager import TransferManagerShmTEProcess total_clients = self.model_config.total_clients if self.model_config.local_dp_size is not None: diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 0a27c850d..e8e5758e1 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -1,42 +1,38 @@ # SPDX-License-Identifier: Apache-2.0 # cython: boundscheck=True, wraparound=True """ -Bootstrap for the radixshmem-backed CPU tier. - -One ``radix-server`` process per node owns the radix index shm, the SlotStore -(the CPU KV pool: one slot per block, a FULL pool and optionally an SWA pool), -the RDMA transfer engine and the etcd data-plane entry. In this mode FlexKV -neither allocates CPU KV memory nor moves bytes between nodes itself: - - * the bootstrap DP process (instance 0, dp 0) launches the server from the - FlexKV configuration -- geometry from ``CacheConfig``, everything else from - the YAML at ``FLEXKV_RADIXSHMEM_CONFIG_PATH`` (``flexkv.common.radixshmem_config``) - -- (``FLEXKV_RADIX_SERVER_LAUNCH_MODE=embedded``) or expects one started by - the operator (``external``); - * every DP scheduler process, the TE process and its transfer workers attach - with ``shmradix.RadixClient(name)``: index operations, ``store`` (the - SlotStore mapping) and ``pull_async`` (the server-side peer pull). - -Naming: index ``/shmradix__cpu`` where ``local_id`` is the YAML's -``cluster.cluster_id`` (plus ``_`` when FLEXKV_RADIX_NODE_NAME names -one of several co-located nodes). A cluster node's index gets ``_`` -appended by radixshmem itself and is resolved through the gRPC socket -``/dev/shm/shmradix__cpu.sock``, so attachers only need the base -name. The SlotStore is ``_data``. - -Geometry: FlexKV stays the source of slot counts and slot bytes. One FULL slot -holds exactly one CPU block as ``StorageEngine`` lays it out (BLOCKFIRST: all -layers of a block contiguous), one SWA slot one SWA page. ``slot_align`` is -chosen so that the SlotStore stride equals the block size exactly, which lets -the H2D / D2H workers address the pool with the strides of a plain tensor. The -TE re-checks the attached regions against the layouts it builds (``check_geometry``). +Attaching the radixshmem-backed CPU tier to this node's radix-server. + +The radix-server is a process the operator starts on every node, with nothing +model-specific on its command line:: + + radix-server --name /flexkv --data-bytes 64G [--swa-ratio 0.5] [cluster flags] + +It owns the radix index shm, the SlotStore (the CPU KV pool), the RDMA transfer +engine and the etcd membership. FlexKV neither creates it nor sizes it: + + * the first FlexKV client hands the server FlexKV's *geometry* -- tokens per + block, bytes of one CPU block, bytes of one SWA page and the SWA window, + the slot alignment that keeps the SlotStore stride equal to the block + (``RadixGeometry``, ``expected_geometry``); the server plans the slot + COUNTS from its byte budget and publishes them (idempotent: every further + client with the same geometry is accepted, one with another geometry is + refused); + * every FlexKV process (the DP schedulers, the TE and its workers) attaches by + the server's name (``attach_radix_client``), checks the published regions + against its own layout (``check_geometry``) and takes the slot counts over + into ``CacheConfig`` (``adopt_geometry``): ``num_cpu_blocks`` and + ``swa.num_slots`` are the server's, not ``cpu_cache_gb``'s. + +Geometry: one FULL slot holds exactly one CPU block as ``StorageEngine`` lays +it out (BLOCKFIRST: all layers of a block contiguous), one SWA slot one SWA +page. ``slot_align`` is chosen so that the SlotStore stride equals the block +size exactly, which lets the H2D / D2H workers address the pool with the +strides of a plain tensor. """ from __future__ import annotations import dataclasses -import multiprocessing as mp -import os -import signal import time from typing import Any, Dict, List, Optional @@ -45,8 +41,7 @@ from flexkv.common.config import (GLOBAL_CONFIG_FROM_ENV, CacheConfig, LayerGroupSpec, ModelConfig, SWAPoolConfig) from flexkv.common.debug import flexkv_logger -from flexkv.common.radixshmem_config import (REGISTER_CHUNK_TOKENS, RadixShmemConfig, - get_radixshmem_config) +from flexkv.common.radixshmem_config import RadixShmemConfig, get_radixshmem_config from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType try: @@ -55,10 +50,8 @@ shmradix = None -_SHM_PREFIX = "/shmradix" # Pool bases are page aligned regardless; a larger per-slot alignment only pads. _MAX_SLOT_ALIGN = 4096 -DEFAULT_HUGETLBFS_DIR = "/mnt/hugepages" def _ensure_shmradix() -> None: @@ -66,26 +59,11 @@ def _ensure_shmradix() -> None: raise ImportError( "shmradix is not installed; install it from the radixshmem repo " "(pip install -e radixshmem/python)") - for name in ("RadixServer", "RadixServerConfig", "IndexConfig", "DataPlaneConfig", - "ClusterConfig", "RadixClient"): + for name in ("RadixClient", "Geometry", "GeometryMismatch", "ServerNotReady"): if not hasattr(shmradix, name): raise ImportError( - f"shmradix lacks {name}: FlexKV needs the RadixServer / RadixClient " - f"surface of radixshmem (transfer-server branch or later)") - - -def radix_index_name(local_id: str) -> str: - """Base shm name of the CPU tier's index for ``local_id`` - (``RadixShmemConfig.local_id``); it also names the gRPC socket.""" - return f"{_SHM_PREFIX}_{local_id}_cpu" - - -def radix_data_name(local_id: str) -> str: - return radix_index_name(local_id) + "_data" - - -def radix_socket_path(local_id: str) -> str: - return "/dev/shm/" + radix_index_name(local_id).lstrip("/").replace("/", "_") + ".sock" + f"shmradix lacks {name}: FlexKV needs a radixshmem whose radix-server takes " + f"its geometry from the client (RadixClient(name, Geometry))") # ------------------------------------------------------------------ geometry @@ -135,8 +113,10 @@ def cpu_block_bytes(model_config: ModelConfig, cache_config: CacheConfig) -> int def swa_pool_config(cache_config: CacheConfig) -> Optional[SWAPoolConfig]: + """The SWA tier FlexKV wants on the server, or None. The slot count is the + server's business (it may still be the placeholder ``cpu_cache_gb`` gave).""" swa = cache_config.swa - if swa is None or not swa.enabled or swa.num_slots <= 0: + if swa is None or not swa.enabled: return None return swa @@ -178,29 +158,38 @@ def slot_align_for(*sizes: int) -> int: @dataclasses.dataclass(frozen=True) class RadixGeometry: - """What FlexKV expects the server's regions to look like.""" + """FlexKV's side of the geometry: what one slot of each pool must hold. The + slot counts are not here; the server plans them from its byte budget.""" tokens_per_block: int - full_slots: int full_slot_bytes: int - swa_slots: int = 0 swa_slot_bytes: int = 0 swa_window_blocks: int = 0 + @property + def has_swa(self) -> bool: + return self.swa_slot_bytes > 0 + @property def slot_align(self) -> int: return slot_align_for(self.full_slot_bytes, self.swa_slot_bytes) - @property - def data_bytes(self) -> int: - return self.full_slots * self.full_slot_bytes + self.swa_slots * self.swa_slot_bytes + def to_shmradix(self) -> "shmradix.Geometry": + """The ``shmradix.Geometry`` handed to the server (data mode: bytes and + window only, no counts).""" + _ensure_shmradix() + return shmradix.Geometry( + block_size=int(self.tokens_per_block), + full_slot_bytes=int(self.full_slot_bytes), + swa_slot_bytes=int(self.swa_slot_bytes), + swa_window_blocks=int(self.swa_window_blocks) if self.has_swa else 0, + slot_align=int(self.slot_align), + ) def describe(self) -> str: - s = (f"tokens_per_block={self.tokens_per_block}, FULL {self.full_slots} x " - f"{self.full_slot_bytes} B") - if self.swa_slots: - s += (f", SWA {self.swa_slots} x {self.swa_slot_bytes} B " - f"(window {self.swa_window_blocks})") - return s + f", slot_align={self.slot_align}, data={self.data_bytes / 2**30:.2f} GiB" + s = f"tokens_per_block={self.tokens_per_block}, FULL slot {self.full_slot_bytes} B" + if self.has_swa: + s += f", SWA slot {self.swa_slot_bytes} B (window {self.swa_window_blocks})" + return s + f", slot_align={self.slot_align}" def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> RadixGeometry: @@ -208,11 +197,8 @@ def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> R raise ValueError( "radixshmem needs FLEXKV_CPU_LAYOUT=BLOCKFIRST: one SlotStore slot is one " "contiguous block, which LAYERFIRST does not give") - if cache_config.num_cpu_blocks <= 0: - raise ValueError(f"cache_config.num_cpu_blocks={cache_config.num_cpu_blocks} must be > 0") geo = RadixGeometry( - tokens_per_block=cache_config.tokens_per_block, - full_slots=int(cache_config.num_cpu_blocks), + tokens_per_block=int(cache_config.tokens_per_block), full_slot_bytes=cpu_block_bytes(model_config, cache_config), ) swa = swa_pool_config(cache_config) @@ -220,49 +206,152 @@ def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> R if swa.window_blocks < 1: raise ValueError( f"cache_config.swa.window_blocks={swa.window_blocks} must be >= 1") - if swa.num_slots < swa.window_blocks: - # All-or-none window allocation: a pool smaller than one window can - # never store anything, so fail at startup. - raise ValueError( - f"cache_config.swa.num_slots={swa.num_slots} cannot hold one " - f"{swa.window_blocks}-block SWA window; raise num_slots or disable SWA") geo = dataclasses.replace( - geo, swa_slots=int(swa.num_slots), - swa_slot_bytes=swa_block_bytes(model_config, cache_config), + geo, swa_slot_bytes=swa_block_bytes(model_config, cache_config), swa_window_blocks=int(swa.window_blocks)) return geo +# -------------------------------------------------------------------- attach + +def attach_radix_client(name: Optional[str] = None, + *, + geometry: Any = None, + rcfg: Optional[RadixShmemConfig] = None, + endpoint: Optional[str] = None, + timeout_s: Optional[float] = None, + max_outstanding: Optional[int] = None, + label: str = "radixshmem") -> "shmradix.RadixClient": + """A ready ``shmradix.RadixClient`` on the radix-server ``name`` (default: + the configuration's ``server.name``). + + With ``geometry`` (a :class:`RadixGeometry` or a ``shmradix.Geometry``) the + client hands the server FlexKV's slot shape on the way; the server plans + the counts from its budget. That is idempotent, so every FlexKV process + may bring it; a server already serving another geometry (another model or + page size on this node) is refused, and so is one whose budget cannot hold + FlexKV's slots. Without a geometry the call only waits for a server that + somebody else configured. + + Retries while the server is not reachable yet (the operator may start it + late), then blocks in ``wait_ready`` -- the rendezvous of a cluster and + the SlotStore prefault happen there -- for ``timeout_s`` in total + (default: the configuration's ``server.ready_timeout_s``). + """ + _ensure_shmradix() + if rcfg is None: + rcfg = get_radixshmem_config() + name = name or rcfg.server_name + if endpoint is None: + endpoint = rcfg.endpoint or None + if timeout_s is None: + timeout_s = rcfg.ready_timeout_s + if max_outstanding is None: + max_outstanding = rcfg.client.max_outstanding + spec = geometry.to_shmradix() if isinstance(geometry, RadixGeometry) else geometry + where = endpoint or f"unix:///dev/shm/{name.lstrip('/').replace('/', '_')}.sock" + + deadline = time.monotonic() + float(timeout_s) + last: Optional[BaseException] = None + while True: + try: + client = shmradix.RadixClient(name, spec, endpoint=endpoint, + max_outstanding=max_outstanding) + break + except shmradix.GeometryMismatch as e: + raise ValueError( + f"{label}: radix-server {name} already serves another geometry ({e}); every " + f"engine attached to one server must run the same model, page size and SWA " + f"configuration") from e + except ValueError as e: + raise ValueError( + f"{label}: radix-server {name} cannot serve FlexKV's geometry ({e}); check its " + f"--data-bytes / --swa-ratio") from e + except Exception as e: # noqa: BLE001 - not reachable yet: no socket, no listener + last = e + if time.monotonic() >= deadline: + raise TimeoutError( + f"{label}: no radix-server named {name} reachable at {where} within " + f"{timeout_s:.0f}s (start it with `radix-server --name {name} " + f"--data-bytes ...`): {last}") from e + flexkv_logger.debug(f"{label}: radix-server {name} not reachable yet ({e}); retrying") + time.sleep(0.5) + + remaining = max(1.0, deadline - time.monotonic()) + try: + info = client.wait_ready(remaining) + except TimeoutError as e: + mode, err = client.info.mode, client.info.last_error + client.close() + raise TimeoutError( + f"{label}: radix-server {name} not ready within {timeout_s:.0f}s (mode={mode}" + + (f", last_error={err!r}" if err else "") + + ("; nobody handed it a geometry" if spec is None and mode == "waiting" else "") + + ")") from e + except RuntimeError as e: + client.close() + raise RuntimeError(f"{label}: radix-server {name} failed to configure: {e}") from e + flexkv_logger.info( + f"{label}: attached radix-server {name} ({where}): index={info.index_name}, " + f"rank={info.rank}/{info.world_size}, data_plane={info.data_plane}, " + f"geometry={_describe_published(info.geometry)}") + return client + + +def _describe_published(g: Optional[Dict[str, Any]]) -> str: + if not g: + return "(none)" + parts = [f"block_size={g.get('block_size')}"] + for kind, pool in (g.get("pools") or {}).items(): + s = f"{kind.upper()} {pool.get('num_slots')} x {pool.get('slot_bytes')} B" + if kind == "swa": + s += f" (window {pool.get('window_blocks')})" + parts.append(s) + return ", ".join(parts) + + +def _published_geometry(client: "shmradix.RadixClient", label: str) -> Dict[str, Any]: + g = client.geometry + if not g: + info = client.status() + g = info.geometry + if not g: + raise ValueError( + f"{label}: radix-server {client.name} has no geometry yet (mode={info.mode}); " + f"attach with FlexKV's geometry first") + return g + + def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, label: str = "radixshmem") -> None: - """Fail closed when the attached regions differ from FlexKV's own layout: a - slot count or stride mismatch would otherwise become a silent misaddressed - transfer.""" - g = client.geometry + """Fail closed when the server's regions differ from FlexKV's own layout: a + stride or page mismatch would otherwise become a silent misaddressed + transfer. Slot counts are not checked here -- they are the server's, taken + over by :func:`adopt_geometry`.""" + _ensure_shmradix() + g = _published_geometry(client, label) pools = g["pools"] diffs: List[str] = [] if int(g["block_size"]) != expected.tokens_per_block: diffs.append(f"tokens_per_block server={g['block_size']} flexkv={expected.tokens_per_block}") full = pools["full"] - if int(full["num_slots"]) != expected.full_slots: - diffs.append(f"FULL slots server={full['num_slots']} flexkv={expected.full_slots}") if int(full["slot_bytes"]) != expected.full_slot_bytes: diffs.append(f"FULL slot_bytes server={full['slot_bytes']} flexkv={expected.full_slot_bytes}") if not client.info.data_plane: - diffs.append("server is index-only (no SlotStore); FlexKV needs the data plane") + diffs.append("server is index-only (no --data-bytes); FlexKV needs the data plane") else: store = client.store stride = int(store.pool(shmradix.ComponentType.FULL).slot_bytes) if stride != expected.full_slot_bytes: diffs.append(f"FULL stride server={stride} flexkv={expected.full_slot_bytes} " - f"(slot_align must divide the block size)") + f"(the server rounds slots up to slot_align={g.get('slot_align')}; " + f"FlexKV asks for {expected.slot_align})") swa = pools.get("swa") - if expected.swa_slots > 0: + if expected.has_swa: if swa is None: - diffs.append("server has no SWA pool but FlexKV's SWA tier is on") + diffs.append("server has no SWA pool but FlexKV's SWA tier is on " + "(start it with --swa-ratio)") else: - if int(swa["num_slots"]) != expected.swa_slots: - diffs.append(f"SWA slots server={swa['num_slots']} flexkv={expected.swa_slots}") if int(swa["slot_bytes"]) != expected.swa_slot_bytes: diffs.append(f"SWA slot_bytes server={swa['slot_bytes']} flexkv={expected.swa_slot_bytes}") if int(swa.get("window_blocks", 0)) != expected.swa_window_blocks: @@ -276,208 +365,56 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, diffs.append("server has an SWA pool that FlexKV's configuration does not") if diffs: raise ValueError( - f"{label}: the attached radixshmem regions do not match FlexKV's configuration " + f"{label}: the radix-server's regions do not match FlexKV's layout " f"({expected.describe()}): " + "; ".join(diffs)) -# ------------------------------------------------------------- server config - -def build_radix_server_config(model_config: ModelConfig, - cache_config: CacheConfig, - rcfg: Optional[RadixShmemConfig] = None, - ) -> "shmradix.RadixServerConfig": - """The one radix-server this FlexKV node needs: index sized from - ``cache_config`` (FULL slots = ``num_cpu_blocks``, SWA slots = ``swa.num_slots``), - SlotStore sized so that every slot is exactly one block, and the cluster / - data-plane / index / server settings of ``rcfg`` (default: the process's - ``FLEXKV_RADIXSHMEM_CONFIG_PATH``) passed through. Raises on an - inconsistent configuration.""" - _ensure_shmradix() - if rcfg is None: - rcfg = get_radixshmem_config() - geo = expected_geometry(model_config, cache_config) - index_kwargs = dict(rcfg.index) - # an RHT registration chunk covers REGISTER_CHUNK_TOKENS tokens unless the file says otherwise - index_kwargs.setdefault("register_chunk_size", - max(1, REGISTER_CHUNK_TOKENS // geo.tokens_per_block)) - index = shmradix.IndexConfig( - name=radix_index_name(rcfg.local_id), - tokens_per_block=geo.tokens_per_block, - full_slots=geo.full_slots, - swa_slots=geo.swa_slots, - swa_window_blocks=geo.swa_window_blocks, - **index_kwargs, - ) - data = shmradix.DataPlaneConfig( - data_bytes=geo.data_bytes, - full_slot_bytes=geo.full_slot_bytes, - swa_slot_bytes=geo.swa_slot_bytes, - slot_align=geo.slot_align, - data_name=radix_data_name(rcfg.local_id), - **rcfg.data, - ) - cluster = shmradix.ClusterConfig(**rcfg.cluster) - server_kwargs = dict(rcfg.server) - if not server_kwargs.get("hugepage_path") and cache_config.use_hugepage_cpu_buffer: - server_kwargs["hugepage_path"] = os.environ.get("FLEXKV_HUGETLBFS_DIR", - DEFAULT_HUGETLBFS_DIR) - cfg = shmradix.RadixServerConfig(index=index, data=data, cluster=cluster, **server_kwargs) - flexkv_logger.info( - f"radixshmem server config for {index.name}: {geo.describe()}, " - f"{rcfg.describe()}, register_chunk_size={index.register_chunk_size}, " - f"hugepage_path={cfg.hugepage_path or '(shm)'}, " - f"prefault={data.prefault}, transfer_devices={data.transfer_devices or '(all)'}") - return cfg - - -# ------------------------------------------------------------ server process - -def _radix_server_main(cfg, ready, stop, conn) -> None: - """Body of the radix-server subprocess: bring the server up, report, wait.""" - import shmradix as _shmradix - - def _on_term(signum, frame): # noqa: ARG001 - stop.set() - - signal.signal(signal.SIGTERM, _on_term) - signal.signal(signal.SIGINT, _on_term) - try: - server = _shmradix.RadixServer(cfg) - server.start() - except BaseException as e: # noqa: BLE001 - reported to the parent, which raises - try: - conn.send(("error", f"{type(e).__name__}: {e}")) - finally: - conn.close() - return - index = server.index - try: - conn.send(("ready", { - "index_name": index.shm_name(), - "rank": int(index.rank()), - "world_size": int(index.world_size()), - "distributed": bool(index.is_distributed()), - })) - finally: - conn.close() - ready.set() +def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", + label: str = "radixshmem") -> Dict[str, int]: + """Take the slot counts the server planned from its budget over into + ``cache_config``: ``num_cpu_blocks`` = the FULL pool, ``swa.num_slots`` = + the SWA pool. Whatever ``cpu_cache_gb`` had produced was a placeholder in + this mode. Run it after any ``recompute_cache_block_counts`` in the same + process (that recompute sizes from ``cpu_cache_gb`` and would undo this). + Returns ``{"full": n, "swa": m}``.""" + g = _published_geometry(client, label) + pools = g["pools"] + counts: Dict[str, int] = {"full": int(pools["full"]["num_slots"])} + before = int(cache_config.num_cpu_blocks) + cache_config.num_cpu_blocks = counts["full"] + note = f"FULL {counts['full']} slots (cpu_cache_gb had given {before})" + swa = swa_pool_config(cache_config) + if swa is not None: + published = pools.get("swa") + if published is None: + raise ValueError( + f"{label}: FlexKV's SWA tier is on but radix-server {client.name} planned no SWA " + f"pool; start it with --swa-ratio") + counts["swa"] = int(published["num_slots"]) + swa_before = int(swa.num_slots) + swa.num_slots = counts["swa"] + note += f", SWA {counts['swa']} slots (had {swa_before})" + flexkv_logger.info(f"{label}: adopted radix-server {client.name}'s slot counts: {note}") + return counts + + +def radix_server_is_distributed(rcfg: Optional[RadixShmemConfig] = None, + *, + timeout_s: Optional[float] = None, + label: str = "radixshmem") -> bool: + """Whether this node's radix-server is part of a cluster (world_size > 1), + i.e. whether peer pulls are possible. Asked of the server itself once it + is ready (a geometry-less attach that only waits), so every process that + asks gets the same answer -- the framework adapters gate the prefetch path + on it in every TP rank, and the ranks must agree.""" + client = attach_radix_client(rcfg=rcfg, timeout_s=timeout_s, label=label) try: - while not stop.wait(0.5): - pass - except KeyboardInterrupt: - pass + return int(client.info.world_size) > 1 finally: - server.close() - - -class RadixServerProcess: - """The embedded radix-server: a spawned subprocess running ``RadixServer``. - - Not the scheduler process (its gRPC threads and transfer polling thread - would contend for the GIL, and a clustered server holds RDMA contexts that - do not survive a fork) and not the TE process (which cannot come up before - the GPU registrations, while the CEs attach the index at construction). - """ - - def __init__(self, cfg: "shmradix.RadixServerConfig"): - self.cfg = cfg - self._ctx = mp.get_context("spawn") - self._ready = self._ctx.Event() - self._stop = self._ctx.Event() - self.process = None - self.info: Dict[str, Any] = {} - - def start(self, timeout_s: Optional[float] = None) -> "RadixServerProcess": - if timeout_s is None: - timeout_s = float(self.cfg.cluster.bootstrap_timeout_sec) + 60.0 - parent, child = self._ctx.Pipe(duplex=False) - self.process = self._ctx.Process( - target=_radix_server_main, - args=(self.cfg, self._ready, self._stop, child), - name="flexkv-radix-server", - daemon=True, - ) - self.process.start() - child.close() - deadline = time.monotonic() + timeout_s - try: - while True: - if parent.poll(0.2): - kind, payload = parent.recv() - break - if not self.process.is_alive(): - raise RuntimeError("radix-server exited during startup (see its log)") - if time.monotonic() > deadline: - self.shutdown() - raise TimeoutError( - f"radix-server did not become ready within {timeout_s:.0f}s " - f"(cluster rendezvous or SlotStore prefault still pending?)") - finally: - parent.close() - if kind == "error": - self.shutdown() - raise RuntimeError(f"radix-server failed to start: {payload}") - self.info = payload - flexkv_logger.info( - f"radix-server pid={self.process.pid} ready: index={payload['index_name']} " - f"rank={payload['rank']}/{payload['world_size']} distributed={payload['distributed']}") - return self - - @property - def cluster_rank(self) -> int: - return int(self.info.get("rank", 0)) - - def shutdown(self, timeout: float = 15.0) -> None: - if self.process is None: - return - self._stop.set() - self.process.join(timeout) - if self.process.is_alive(): - self.process.terminate() - self.process.join(5.0) - self.process = None - - -# ------------------------------------------------------------------- attach - -def attach_radix_client(name: str, - timeout_s: Optional[float] = None, - *, - rcfg: Optional[RadixShmemConfig] = None, - max_outstanding: Optional[int] = None) -> "shmradix.RadixClient": - """``shmradix.RadixClient(name)``, retried until the server's socket exists - and the server is ready (an embedded server starts concurrently with the - CEs, an external one may still be rendezvousing). Endpoint, timeout and - ``max_outstanding`` default to ``rcfg`` (the process's configuration). - - The returned client owns the index attach, the SlotStore mapping and the - gRPC channel; keep it alive for as long as its slots are addressed. - """ - _ensure_shmradix() - if rcfg is None: - rcfg = get_radixshmem_config() - if timeout_s is None: - timeout_s = rcfg.attach_timeout_s - if max_outstanding is None: - max_outstanding = rcfg.client.max_outstanding - deadline = time.monotonic() + timeout_s - last: Optional[BaseException] = None - while True: - remaining = deadline - time.monotonic() - if remaining <= 0: - raise TimeoutError( - f"radix-server {name} not attachable within {timeout_s:.0f}s: {last}") - try: - return shmradix.RadixClient(name, endpoint=rcfg.endpoint or None, - timeout_s=max(1.0, remaining), - max_outstanding=max_outstanding) - except Exception as e: # noqa: BLE001 - socket not there yet, server starting - last = e - flexkv_logger.debug(f"attach to radix-server {name} failed (will retry): {e}") - time.sleep(0.2) + client.close() def radix_cluster_rank(client: "shmradix.RadixClient") -> int: - """The cluster rank etcd assigned this node (0 when standalone).""" + """This node's rank in the radix cluster (0 on a standalone server).""" rank = int(getattr(client.info, "rank", -1)) return rank if rank >= 0 else int(client.rank()) diff --git a/flexkv/transfer_manager.py b/flexkv/transfer_manager.py index 0bcee9aa3..62bf21b32 100644 --- a/flexkv/transfer_manager.py +++ b/flexkv/transfer_manager.py @@ -431,11 +431,17 @@ def initialize_transfer_engine(self) -> None: radix_client = None if GLOBAL_CONFIG_FROM_ENV.enable_radixshmem: - from flexkv.common.radixshmem_config import get_radixshmem_config - from flexkv.server.shm_radix_bootstrap import (attach_radix_client, - radix_index_name) - radix_client = attach_radix_client( - radix_index_name(get_radixshmem_config().local_id)) + # The operator's radix-server; its SlotStore is this TE's CPU pool. + # Its slot counts are authoritative: take them over again here, + # AFTER the recompute above (which sizes from cpu_cache_gb and would + # otherwise undo what the KVManager adopted), so the CPU layouts + # below match the server's pools exactly. + from flexkv.server.shm_radix_bootstrap import (adopt_geometry, attach_radix_client, + check_geometry, expected_geometry) + geometry = expected_geometry(self.model_config, self.cache_config) + radix_client = attach_radix_client(geometry=geometry, label="TransferManager") + check_geometry(radix_client, geometry, label="TransferManager") + adopt_geometry(self.cache_config, radix_client, label="TransferManager") self._radix_client = radix_client self.storage_engine = StorageEngine( self.model_config, diff --git a/tests/radixshmem/radix_e2e_common.py b/tests/radixshmem/radix_e2e_common.py index ca77b9a98..deb7f68b6 100644 --- a/tests/radixshmem/radix_e2e_common.py +++ b/tests/radixshmem/radix_e2e_common.py @@ -19,6 +19,7 @@ import shutil import socket import subprocess +import sys import tempfile import time from typing import List, Optional, Tuple @@ -101,6 +102,44 @@ def write_radix_config(workdir: str, config: dict, name: str = "radixshmem.yaml" return path +def start_radix_server(name: str, data_bytes: int, *, extra_args=(), endpoint: Optional[str] = None, + log_path: Optional[str] = None, timeout: float = 60.0) -> subprocess.Popen: + """Start the operator's ``radix-server`` (``python -m shmradix.cli``) and wait + for its socket. Nothing model-specific goes on its command line: FlexKV's + clients bring the geometry, the server plans the slot counts from + ``data_bytes``.""" + cmd = [sys.executable, "-m", "shmradix.cli", "--name", name, "--data-bytes", str(int(data_bytes)), + "--no-prefault", "--interval", "0", *extra_args] + if endpoint: + cmd += ["--endpoint", endpoint] + log = open(log_path, "w") if log_path else subprocess.DEVNULL + proc = subprocess.Popen(cmd, stdout=log, stderr=subprocess.STDOUT) + if endpoint and endpoint.startswith("unix://"): + sock = endpoint[len("unix://"):] + else: + sock = f"/dev/shm/{name.lstrip('/').replace('/', '_')}.sock" + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if proc.poll() is not None: + raise RuntimeError(f"radix-server {name} exited with {proc.returncode}" + + (f"; see {log_path}" if log_path else "")) + if os.path.exists(sock): + return proc + time.sleep(0.2) + proc.terminate() + raise TimeoutError(f"radix-server {name} did not open {sock} within {timeout:.0f}s") + + +def stop_radix_server(proc: Optional[subprocess.Popen]) -> None: + if proc is None: + return + proc.terminate() + with contextlib.suppress(Exception): + proc.wait(20) + if proc.poll() is None: + proc.kill() + + def sweep_radix_files(cluster_id: str) -> None: """Drop what a run under ``cluster_id`` (the radixshmem namespace; every shm / socket / IPC name of the run contains it) may have left in shm / tmp.""" diff --git a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py index bd55a53d3..1d382f26f 100644 --- a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -4,12 +4,13 @@ Two full FlexKV nodes on one host (two processes, two GPUs), one radixshmem cluster: - * each node's KVManager launches its own radix-server (index + SlotStore = - the node's CPU pool + RDMA transfer engine) from one shared YAML - (FLEXKV_RADIXSHMEM_CONFIG_PATH: cluster_id, expected_min_nodes=2, registry, - RDMA devices), told apart by the per-node FLEXKV_RADIX_NODE_NAME override - as co-located nodes are; the two servers rendezvous in one etcd namespace, - get dense cluster ranks and an RHT to route by; + * the test starts one operator-style radix-server per node (index + + SlotStore = the node's CPU pool + RDMA transfer engine), joined into one + cluster by their command lines (--expected-min-nodes 2, --registry, + --node-name, RDMA devices); each FlexKV node attaches to its own server + through a per-node YAML (FLEXKV_RADIXSHMEM_CONFIG_PATH: server.name / + server.endpoint) and brings the geometry; the two servers rendezvous in + one etcd namespace, get dense cluster ranks and an RHT to route by; * node 0 PUTs a window of GPU blocks holding a per-block pattern; * node 1 calls ``KVManager.prefetch_async`` for the same tokens: the index walk finds the prefix on node 0 over RDMA, node 1's radix-server RDMA-reads the @@ -54,6 +55,8 @@ start_tp_client, stop_private_etcd, stop_tp_client, + start_radix_server, + stop_radix_server, sweep_radix_files, wait_kv_manager_ready, write_pattern, @@ -115,12 +118,7 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, recv_port = f"ipc:///tmp/flexkv_{cluster_id}_{node_name}" os.environ.update({ "FLEXKV_ENABLE_RADIXSHMEM": "1", - "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, - # The two per-node overrides of the global file: etcd keys membership - # by node identity, which defaults to the bind IP the co-located nodes - # share -- so name each node and give the loopback address explicitly. - "FLEXKV_RADIX_NODE_NAME": node_name, - "FLEXKV_RADIX_RPC_ADDRESS": "127.0.0.1", + "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, # this node's server "FLEXKV_ENABLE_MPS": "0", "FLEXKV_SERVER_RECV_PORT": recv_port, }) @@ -132,8 +130,6 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, # case a parent import happened earlier in this process. GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path = config_path - GLOBAL_CONFIG_FROM_ENV.radix_node_name = node_name - GLOBAL_CONFIG_FROM_ENV.radix_rpc_address = "127.0.0.1" GLOBAL_CONFIG_FROM_ENV.enable_mps = False GLOBAL_CONFIG_FROM_ENV.server_recv_port = recv_port @@ -234,16 +230,32 @@ def _run(registry: str, rdma_dev: str) -> dict: # One etcd namespace (and shm prefix) per run keeps concurrent runs apart. cluster_id = f"p2p{os.getpid()}" workdir = tempfile.mkdtemp(prefix="flexkv_radix_p2p_") - config_path = write_radix_config(workdir, { - "cluster": { - "cluster_id": cluster_id, - "expected_min_nodes": WORLD_SIZE, - "registry": registry, - "index_dev": rdma_dev, - "rht_slots_per_bucket": 4, - }, - "data": {"transfer_devices": [rdma_dev], "prefault": False}, - }) + # Two operator-style radix-servers on one host, one cluster: distinct names, + # sockets and node names, the loopback address as the bootstrap IP. The + # FlexKV nodes bring the geometry (node 1 adopts what node 0 published). + from flexkv.common.config import CacheConfig, ModelConfig + from flexkv.server.shm_radix_bootstrap import cpu_block_bytes + block_bytes = cpu_block_bytes( + ModelConfig(num_layers=2, num_kv_heads=4, head_size=128, dtype=torch.float16, + tp_size=1, dp_size=1), + CacheConfig(tokens_per_block=TOKENS_PER_BLOCK, enable_cpu=True, enable_ssd=False, + num_cpu_blocks=NUM_CPU_BLOCKS)) + sweep_radix_files(cluster_id) + servers, config_paths = [], [] + for rank in range(WORLD_SIZE): + name = f"/{cluster_id}_{_node_name(rank)}" + endpoint = f"unix:///dev/shm/{cluster_id}_{_node_name(rank)}.sock" + servers.append(start_radix_server( + name, NUM_CPU_BLOCKS * block_bytes, endpoint=endpoint, + extra_args=["--expected-min-nodes", str(WORLD_SIZE), "--registry", registry, + "--cluster-id", cluster_id, "--node-name", _node_name(rank), + "--rpc-address", "127.0.0.1", "--index-dev", rdma_dev, + "--transfer-dev", rdma_dev, "--rht-slots", "4", + "--bootstrap-timeout", "120"], + log_path=os.path.join(workdir, f"radix-server-{_node_name(rank)}.log"))) + config_paths.append(write_radix_config( + workdir, {"server": {"name": name, "endpoint": endpoint, "ready_timeout_s": 300}}, + name=f"radixshmem_{_node_name(rank)}.yaml")) ctx = mp.get_context("spawn") reader_ready, written, read_done = ctx.Event(), ctx.Event(), ctx.Event() result_q = ctx.Queue() @@ -253,7 +265,7 @@ def _run(registry: str, rdma_dev: str) -> dict: for rank in range(WORLD_SIZE): proc = ctx.Process( target=_node_proc, - args=(rank, rank, cluster_id, config_path, + args=(rank, rank, cluster_id, config_paths[rank], reader_ready, written, read_done, result_q), daemon=False, ) @@ -273,6 +285,8 @@ def _run(registry: str, rdma_dev: str) -> dict: if proc.is_alive(): proc.terminate() proc.join(timeout=10) + for server in servers: + stop_radix_server(server) sweep_radix_files(cluster_id) shutil.rmtree(workdir, ignore_errors=True) return reports diff --git a/tests/radixshmem/test_e2e_radix_shmem.py b/tests/radixshmem/test_e2e_radix_shmem.py index 84883c683..70074f27f 100644 --- a/tests/radixshmem/test_e2e_radix_shmem.py +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -3,13 +3,13 @@ For every ``dp_size`` in the parametrization: - * dp0 is the bootstrap process: its KVManager launches the radix-server - (index + SlotStore, the node's CPU KV pool) and spawns the single TE; every - other DP attaches to both by name and feeds the TE over its own shm - channel with a disjoint graph/op id range. The run's namespace is the - ``cluster.cluster_id`` of a small YAML written per run - (FLEXKV_RADIXSHMEM_CONFIG_PATH), which is how a deployment names its - regions too. + * the test starts the operator's radix-server (index + SlotStore, the + node's CPU KV pool) with nothing but a name and a byte budget; dp0's + KVManager hands it FlexKV's geometry, adopts the slot counts it plans and + spawns the single TE; every other DP attaches to both by name and feeds + the TE over its own shm channel with a disjoint graph/op id range. Which + server to attach to is the ``server.name`` of a small YAML written per run + (FLEXKV_RADIXSHMEM_CONFIG_PATH), which is how a deployment names it too. * Phase 1: every DP PUTs its own requests concurrently through the shared TE. * Phase 2 (dp_size > 1): dp0 PUTs a prefix that dp1 then finds with ``get_match`` -- the shared index is what the radixshmem path exists for. @@ -46,6 +46,8 @@ mismatched_blocks, start_tp_client, stop_tp_client, + start_radix_server, + stop_radix_server, sweep_radix_files, wait_kv_manager_ready, write_pattern, @@ -185,8 +187,20 @@ def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, def _run(dp_size: int) -> dict: server_id = f"e2e{dp_size}dp_{os.getpid()}" workdir = tempfile.mkdtemp(prefix="flexkv_radix_e2e_") - config_path = write_radix_config(workdir, {"cluster": {"cluster_id": server_id}, - "data": {"prefault": False}}) + name = f"/{server_id}" + config_path = write_radix_config(workdir, {"server": {"name": name, "ready_timeout_s": 300}}) + # The operator's server: a byte budget that holds NUM_CPU_BLOCKS blocks of + # this test's model (the DP processes bring the geometry and adopt the count). + from flexkv.common.config import CacheConfig, ModelConfig + from flexkv.server.shm_radix_bootstrap import cpu_block_bytes + block_bytes = cpu_block_bytes( + ModelConfig(num_layers=2, num_kv_heads=4, head_size=128, dtype=torch.float16, + tp_size=1, dp_size=dp_size), + CacheConfig(tokens_per_block=TOKENS_PER_BLOCK, enable_cpu=True, enable_ssd=False, + num_cpu_blocks=NUM_CPU_BLOCKS)) + sweep_radix_files(server_id) + server = start_radix_server(name, NUM_CPU_BLOCKS * block_bytes, + log_path=os.path.join(workdir, "radix-server.log")) ctx = mp.get_context("spawn") barrier = ctx.Barrier(dp_size) result_q = ctx.Queue() @@ -214,6 +228,7 @@ def _run(dp_size: int) -> dict: if proc.is_alive(): proc.terminate() proc.join(timeout=10) + stop_radix_server(server) sweep_radix_files(server_id) shutil.rmtree(workdir, ignore_errors=True) return reports diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index da3336103..5bbb5f3ef 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -5,12 +5,15 @@ Four parts, one file: Part 1 — `CacheEngineRadixShmem` semantics against an in-process - `shmradix.RadixServer` (index + SlotStore): take / insert / match / recycle, - insert-after-transfer publication, lock vs. eviction, SWA windows, and the - fact that a standalone region has no peer to prefetch from. + `shmradix.RadixServer` (index + SlotStore, started the way the operator's + `radix-server` is: a name and a byte budget; FlexKV's client brings the + geometry): take / insert / match / recycle, insert-after-transfer + publication, lock vs. eviction, SWA windows, and the fact that a standalone + region has no peer to prefetch from. Part 1b — the data plane: the SlotStore pool viewed as FlexKV's CPU tensor, - the exact-stride geometry the bootstrap derives from the configuration, and - the embedded radix-server subprocess. + the exact-stride geometry FlexKV derives from its configuration and hands + to the server, the slot counts it adopts back, and a client that arrives + before its server. Part 2 — `GlobalCacheEngine.get()/put()` planning on the radixshmem backend, driven by synthetic matches (no region): the local GET is one H2D, the prefetch plan carries a `pull_async` job, the PUT arms the deferred insert; @@ -32,6 +35,7 @@ import contextlib import copy +import dataclasses import glob import importlib.util import multiprocessing as mp @@ -41,6 +45,7 @@ import subprocess import sys import tempfile +import threading import time import traceback from dataclasses import dataclass @@ -59,7 +64,7 @@ pytest.skip(f"shmradix unusable ({exc}); rebuild the extension", allow_module_level=True) -for _name in ("RadixServer", "RadixServerConfig", "IndexConfig", "DataPlaneConfig"): +for _name in ("RadixServer", "ServerConfig", "Geometry", "RadixClient"): if not hasattr(shmradix, _name): pytest.skip(f"shmradix lacks {_name}: needs the RadixServer/RadixClient surface", allow_module_level=True) @@ -141,20 +146,31 @@ def _sweep_region(name: str, data_name: str | None = None) -> None: os.remove(path) +def _server_budget(blocks: int, swa_slots: int = 0, slot_bytes: int = SLOT_BYTES): + """(data_bytes, swa_ratio) that make a data-mode server plan exactly + ``blocks`` FULL and ``swa_slots`` SWA slots of ``slot_bytes`` each (the + stride equals the slot bytes, see ``slot_align_for``).""" + data_bytes = (blocks + swa_slots) * slot_bytes + swa_ratio = (swa_slots * slot_bytes) / data_bytes if swa_slots else 0.0 + return data_bytes, swa_ratio + + def _server_config(name: str, blocks: int, tokens_per_block: int, swa_slots: int = 0, window_blocks: int = 0, slot_bytes: int = SLOT_BYTES): - """A standalone data-mode server: one slot per block, stride == slot bytes.""" + """A standalone data-mode radix-server the way the operator starts it (a + name and a byte budget, nothing about the model) plus the Geometry + FlexKV's client brings. Budget and ratio are chosen so the counts the + server plans equal the test's ``blocks`` / ``swa_slots``.""" align = bootstrap.slot_align_for(slot_bytes) - return shmradix.RadixServerConfig( - index=shmradix.IndexConfig( - name=name, tokens_per_block=tokens_per_block, full_slots=blocks, - swa_slots=swa_slots, swa_window_blocks=window_blocks), - data=shmradix.DataPlaneConfig( - data_bytes=(blocks + swa_slots) * slot_bytes, full_slot_bytes=slot_bytes, - swa_slot_bytes=slot_bytes if swa_slots else 0, slot_align=align, - prefault=False), - ) + data_bytes, swa_ratio = _server_budget(blocks, swa_slots, slot_bytes) + cfg = shmradix.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=swa_ratio, + slot_align=align, prefault=False) + geo = shmradix.Geometry(block_size=tokens_per_block, full_slot_bytes=slot_bytes, + swa_slot_bytes=slot_bytes if swa_slots else 0, + swa_window_blocks=window_blocks if swa_slots else 0, + slot_align=align) + return cfg, geo class _Env: @@ -165,7 +181,8 @@ def __init__(self) -> None: self._stack = [] def server(self, cfg): - _sweep_region(cfg.index.name, cfg.resolved_data_name) + """Start the server: ``waiting`` until a client brings the geometry.""" + _sweep_region(cfg.name, cfg.resolved_data_name) server = shmradix.RadixServer(cfg).start() self._stack.append(server.close) return server @@ -176,9 +193,9 @@ def engine(self, name: str, **kwargs) -> CacheEngineRadixShmem: return engine def make(self, name: str, blocks: int = 2000, tokens_per_block: int = 4, **engine_kwargs): - cfg = _server_config(name, blocks, tokens_per_block) + cfg, geo = _server_config(name, blocks, tokens_per_block) server = self.server(cfg) - engine = self.engine(name, num_total_blocks=blocks, + engine = self.engine(name, geometry=geo, num_total_blocks=blocks, tokens_per_block=tokens_per_block, **engine_kwargs) return engine, server @@ -188,10 +205,10 @@ def close(self) -> None: self._stack.pop()() -def _radix_config(**cluster): - """The all-defaults radixshmem configuration (standalone, socket derived - from the index name) with ``cluster`` keys changed.""" - return load_radixshmem_config(None).replace_cluster(**cluster) +def _radix_config(**server): + """The all-defaults radixshmem configuration with ``server`` keys changed; + a short ready timeout so a broken test fails instead of waiting.""" + return load_radixshmem_config(None).replace_server(**{"ready_timeout_s": 60.0, **server}) @pytest.fixture @@ -374,11 +391,11 @@ def _make_swa_engine(env, name: str, blocks: int = 2000, swa_slots: int = 64, tokens_per_block: int = 16, window_blocks: int = SWA_W): """A single region carrying the SWA component, and an engine that knows it.""" from flexkv.common.config import SWAPoolConfig - cfg = _server_config(name, blocks, tokens_per_block, - swa_slots=swa_slots, window_blocks=window_blocks) + cfg, geo = _server_config(name, blocks, tokens_per_block, + swa_slots=swa_slots, window_blocks=window_blocks) server = env.server(cfg) engine = env.engine( - name, num_total_blocks=blocks, tokens_per_block=tokens_per_block, + name, geometry=geo, num_total_blocks=blocks, tokens_per_block=tokens_per_block, swa_config=SWAPoolConfig(enabled=True, num_slots=swa_slots, window_blocks=window_blocks)) return engine, server @@ -571,84 +588,94 @@ def _configs(num_cpu_blocks: int = 64, swa_slots: int = 0): def test_expected_geometry_mirrors_the_storage_engine_layout(): """One FULL slot is one CPU block exactly as StorageEngine lays it out: 2 layers x 2 (K,V) x 16 tokens x 4 heads x 64 x fp16 = 32768 B; one SWA - slot is one SWA page: 1 layer x 16 tokens x 64 B.""" + slot is one SWA page: 1 layer x 16 tokens x 64 B. Counts are not part of + it: the server plans them from its budget.""" model_config, cache_config = _configs(num_cpu_blocks=64, swa_slots=16) geo = bootstrap.expected_geometry(model_config, cache_config) - assert geo.tokens_per_block == 16 - assert geo.full_slots == 64 and geo.full_slot_bytes == 32768 - assert geo.swa_slots == 16 and geo.swa_slot_bytes == 1024 and geo.swa_window_blocks == SWA_W + assert geo.tokens_per_block == 16 and geo.full_slot_bytes == 32768 + assert geo.has_swa and geo.swa_slot_bytes == 1024 and geo.swa_window_blocks == SWA_W assert geo.slot_align == 1024 # gcd power of two of 32768 and 1024 - assert geo.data_bytes == 64 * 32768 + 16 * 1024 + spec = geo.to_shmradix().to_dict() + assert spec["block_size"] == 16 and spec["slot_align"] == 1024 + assert spec["pools"]["full"] == {"slot_bytes": 32768, "num_slots": 0} + assert spec["pools"]["swa"] == {"slot_bytes": 1024, "num_slots": 0, "window_blocks": SWA_W} + # Without SWA the pool is absent from what the server is asked for. + model_config, cache_config = _configs(num_cpu_blocks=64) + spec = bootstrap.expected_geometry(model_config, cache_config).to_shmradix().to_dict() + assert set(spec["pools"]) == {"full"} -def test_server_config_and_geometry_check(env): - """`build_radix_server_config` starts a server whose regions pass - `check_geometry`; a different expectation is rejected, not papered over.""" +def test_client_brings_the_geometry_and_adopts_the_counts(env): + """The operator's server knows only its budget; FlexKV's first client hands + it the slot shape, the server plans the counts, `check_geometry` verifies + the regions and `adopt_geometry` takes the counts into CacheConfig. A + different geometry is refused, not papered over.""" model_config, cache_config = _configs(num_cpu_blocks=64, swa_slots=16) - rcfg = _radix_config(cluster_id=f"geo{os.getpid()}") - set_radixshmem_config(rcfg) + geo = bootstrap.expected_geometry(model_config, cache_config) + name = f"/geo{os.getpid()}" + data_bytes = 64 * geo.full_slot_bytes + 16 * geo.swa_slot_bytes + env.server(shmradix.ServerConfig(name=name, data_bytes=data_bytes, + swa_ratio=16 * geo.swa_slot_bytes / data_bytes, + prefault=False)) # slot_align comes with the geometry + client = bootstrap.attach_radix_client(name, geometry=geo, timeout_s=60) try: - cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) - assert cfg.index.name == bootstrap.radix_index_name(rcfg.local_id) - assert cfg.cluster.cluster_id == rcfg.cluster_id - # FlexKV's own defaults, where they differ from radixshmem's. - assert cfg.cluster.rht_slots_per_bucket == 4 - assert cfg.cluster.bootstrap_timeout_sec == 120 - assert cfg.index.data_pool_ratio == 8.0 - # one RHT registration chunk per 4096 tokens, whatever the block size - assert cfg.index.register_chunk_size == 4096 // cache_config.tokens_per_block - assert cfg.data.slot_align == 1024 - assert cfg.index.full_slots == 64 and cfg.index.swa_slots == 16 - env.server(cfg) - client = bootstrap.attach_radix_client(cfg.index.name, timeout_s=30) - try: - geo = bootstrap.expected_geometry(model_config, cache_config) - bootstrap.check_geometry(client, geo, "test") - assert int(client.store.pool(FULL).slot_bytes) == geo.full_slot_bytes - assert int(client.store.pool(_SWA).slot_bytes) == geo.swa_slot_bytes - assert bootstrap.radix_cluster_rank(client) == 0 - cache_config.num_cpu_blocks = 65 - with pytest.raises(ValueError, match="FULL slots"): - bootstrap.check_geometry( - client, bootstrap.expected_geometry(model_config, cache_config), "test") - finally: - client.close() + bootstrap.check_geometry(client, geo, "test") + pools = client.geometry["pools"] + assert pools["full"]["num_slots"] == 64 and pools["swa"]["num_slots"] == 16 + assert int(client.store.pool(FULL).slot_bytes) == geo.full_slot_bytes # exact stride + assert int(client.store.pool(_SWA).slot_bytes) == geo.swa_slot_bytes + # cpu_cache_gb's placeholders give way to the server's counts + cache_config.num_cpu_blocks, cache_config.swa.num_slots = 7, 3 + assert bootstrap.adopt_geometry(cache_config, client, "test") == {"full": 64, "swa": 16} + assert cache_config.num_cpu_blocks == 64 and cache_config.swa.num_slots == 16 + assert bootstrap.radix_cluster_rank(client) == 0 and client.info.world_size == 1 + # the connectors' prefetch gate asks the server the same question + assert bootstrap.radix_server_is_distributed(_radix_config(name=name), timeout_s=30) is False + # another expectation against the same regions fails closed + cache_config.tokens_per_block = 32 + with pytest.raises(ValueError, match="tokens_per_block"): + bootstrap.check_geometry( + client, bootstrap.expected_geometry(model_config, cache_config), "test") + # ...and a second client bringing another geometry is refused by the server + other = dataclasses.replace(geo, tokens_per_block=32) + with pytest.raises(ValueError, match="another geometry"): + bootstrap.attach_radix_client(name, geometry=other, timeout_s=60) finally: - set_radixshmem_config(None) + client.close() -def test_embedded_server_process_lifecycle(): - """The bootstrap DP process runs the radix-server as a spawned subprocess: - start() returns once it is ready, clients attach by name, shutdown() takes - the socket down with it.""" - model_config, cache_config = _configs(num_cpu_blocks=64) - rcfg = _radix_config(cluster_id=f"proc{os.getpid()}") - set_radixshmem_config(rcfg) - shm_radix_id = rcfg.local_id - cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) - _sweep_region(cfg.index.name, cfg.resolved_data_name) - server = bootstrap.RadixServerProcess(cfg) +def test_attach_waits_for_a_late_server(env): + """The operator may start the radix-server after the engine: the attach + retries until the socket answers, brings the geometry and waits for + ready. No server at all fails with the command that starts one.""" + name = f"/late{os.getpid()}" + cfg, geo = _server_config(name, blocks=32, tokens_per_block=4) + _sweep_region(cfg.name, cfg.resolved_data_name) + started = threading.Event() + + def _start_later(): + time.sleep(1.5) + env._stack.append(shmradix.RadixServer(cfg).start().close) + started.set() + + threading.Thread(target=_start_later, daemon=True).start() + t0 = time.monotonic() + client = bootstrap.attach_radix_client(name, geometry=geo, timeout_s=60) try: - server.start(timeout_s=120) - assert server.cluster_rank == 0 - assert server.info["distributed"] is False - assert os.path.exists(bootstrap.radix_socket_path(shm_radix_id)) - client = bootstrap.attach_radix_client(cfg.index.name, timeout_s=30) - assert client.info.data_plane - assert int(client.mempool_total()) == 64 - client.close() + assert started.is_set() and time.monotonic() - t0 >= 1.0 + assert client.info.mode == "ready" and client.info.data_plane + assert int(client.mempool_total()) == 32 + assert bootstrap.radix_cluster_rank(client) == 0 finally: - server.shutdown() - set_radixshmem_config(None) - assert server.process is None - assert not os.path.exists(bootstrap.radix_socket_path(shm_radix_id)) + client.close() + with pytest.raises(TimeoutError, match="radix-server --name"): + bootstrap.attach_radix_client(f"/nobody{os.getpid()}", geometry=geo, timeout_s=2) # ----------------------------------------------------------------------------- -# Part 1c — the radixshmem-mode YAML (flexkv.common.radixshmem_config): pass- -# through sections validated against shmradix's dataclasses, FlexKV's own -# defaults, the per-node overrides, and the startup checks -# (docs/radixshmem/config_zh.md section 6). +# Part 1c — the radixshmem-mode YAML (flexkv.common.radixshmem_config): which +# radix-server to attach to and FlexKV's client settings; the former +# server-side sections are refused (docs/radixshmem/config_zh.md). def _write_yaml(tmp_path, text: str) -> str: @@ -657,89 +684,50 @@ def _write_yaml(tmp_path, text: str) -> str: return str(path) -def test_radix_config_defaults_and_passthrough(tmp_path): +def test_radix_config_defaults(): cfg = load_radixshmem_config(None) - assert cfg.cluster_id == "flexkv" and cfg.local_id == "flexkv" - assert not cfg.distributed and cfg.endpoint == "" - # FlexKV's defaults where they differ from radixshmem's; nothing else is set. - assert cfg.cluster == {"cluster_id": "flexkv", "bootstrap_timeout_sec": 120, - "rht_slots_per_bucket": 4} - assert cfg.index == {"data_pool_ratio": 8.0} and cfg.data == {} and cfg.server == {} - assert cfg.attach_timeout_s == 180.0 + assert cfg.path is None + assert cfg.server_name == "/flexkv" and cfg.endpoint == "" and cfg.ready_timeout_s == 600.0 + assert cfg.te_server_id == "flexkv" + assert cfg.client.prefetch_timeout_ms == 5000 and cfg.client.prefetch_max_inflight == 128 + assert cfg.client.max_outstanding == 256 + assert "radix-server /flexkv" in cfg.describe() + +def test_radix_config_file(tmp_path): path = _write_yaml(tmp_path, """ -cluster: - cluster_id: prod - expected_min_nodes: 3 - registry: etcd://10.0.0.1:2379 - rpc_interface: eth0 - index_dev: mlx5_0 - rht_transport: xrc - peer_index_transport: dc - num_rht_shards: 2 - rht_shard_holders: "0,2" -data: - transfer_devices: mlx5_1,mlx5_2 - prefault: false -index: - background_evict_ratio: 0.1 server: - rpc_workers: 8 + name: /prod/kv + endpoint: 10.0.0.2:7000 + ready_timeout_s: 900 client: prefetch_timeout_ms: 1000 + max_outstanding: 512 + prefetch_max_inflight: 300 """) cfg = load_radixshmem_config(path) - assert cfg.path == path and cfg.distributed and cfg.expected_min_nodes == 3 - assert cfg.cluster["rht_shard_holders"] == [0, 2] - assert cfg.data == {"transfer_devices": ["mlx5_1", "mlx5_2"], "prefault": False} - assert cfg.index == {"data_pool_ratio": 8.0, "background_evict_ratio": 0.1} - assert cfg.server == {"rpc_workers": 8} - assert cfg.client.prefetch_timeout_ms == 1000 and cfg.client.max_outstanding == 256 - # Every section constructs its shmradix dataclass as is. - shmradix.ClusterConfig(**cfg.cluster) - shmradix.IndexConfig(**cfg.index) - shmradix.DataPlaneConfig(data_bytes=1, full_slot_bytes=1, **cfg.data) - - -def test_radix_config_per_node_overrides(tmp_path): - path = _write_yaml(tmp_path, """ -cluster: - cluster_id: prod - expected_min_nodes: 2 - registry: etcd://10.0.0.1:2379 - rpc_interface: eth0 -""") - cfg = load_radixshmem_config(path, node_name="r1", rpc_address="127.0.0.1") - assert cfg.node_name == "r1" and cfg.rpc_address == "127.0.0.1" - # The explicit address must not lose to the file's interface (radixshmem - # lets the interface win), and co-located nodes get distinct regions. - assert cfg.cluster["rpc_interface"] == "" - assert cfg.local_id == "prod_r1" - assert bootstrap.radix_index_name(cfg.local_id) == "/shmradix_prod_r1_cpu" - # Without the interface, the address alone satisfies the cluster check. - path = _write_yaml(tmp_path, "cluster:\n expected_min_nodes: 2\n registry: etcd://h:1\n") - with pytest.raises(RadixShmemConfigError, match="rpc_interface"): - load_radixshmem_config(path) - load_radixshmem_config(path, rpc_address="10.0.0.5") + assert cfg.path == path and cfg.server_name == "/prod/kv" and cfg.te_server_id == "prod_kv" + assert cfg.endpoint == "10.0.0.2:7000" and cfg.ready_timeout_s == 900.0 + assert cfg.client.prefetch_timeout_ms == 1000 and cfg.client.max_outstanding == 512 + assert cfg.client.prefetch_max_inflight == 300 + assert cfg.replace_server(name="/x").server_name == "/x" + assert cfg.replace_client(max_outstanding=1000).client.max_outstanding == 1000 @pytest.mark.parametrize("text, match", [ - ("cluster:\n node_name: n0\n", "per-node"), - ("cluster:\n rpc_address: 10.0.0.1\n", "per-node"), - ("index:\n full_slots: 5\n", "geometry is derived"), - ("data:\n slot_align: 4096\n", "geometry is derived"), - ("cluster:\n transport: xrc\n", "unknown key"), - ("server:\n cluster: {}\n", "unknown key"), + ("cluster:\n cluster_id: prod\n", "radix-server"), # former server-side sections + ("data:\n prefault: false\n", "radix-server"), + ("index:\n data_pool_ratio: 8\n", "radix-server"), ("peers: {}\n", "unknown section"), ("- a\n", "must be a mapping"), - ("cluster:\n expected_min_nodes: 2\n rpc_interface: eth0\n", "registry"), - ("cluster:\n expected_min_nodes: 2\n registry: etcd://h:1\n rpc_interface: eth0\n" - " num_rht_shards: 3\n", "num_rht_shards"), - ("cluster:\n rht_slots_per_bucket: 3\n", "rht_slots_per_bucket"), - ("cluster:\n rht_transport: rc\n", "rht_transport"), - ("cluster:\n peer_index_transport: tcp\n", "peer_index_transport"), - ("cluster:\n remote_op_transport: xrc\n", "remote_op_transport"), + ("server: 5\n", "must be a mapping"), + ("server:\n rpc_workers: 3\n", "unknown key"), + ("server:\n name: kv\n", "starts with"), + ("server:\n name: /a b\n", "starts with"), + ("server:\n ready_timeout_s: 0\n", "ready_timeout_s"), + ("server:\n ready_timeout_s: soon\n", "invalid value"), ("client:\n prefetch_max_inflight: 256\n", "max_outstanding"), + ("client:\n prefetch_timeout_ms: 0\n", "prefetch_timeout_ms"), ("client:\n timeout: 5\n", "unknown key"), ]) def test_radix_config_rejects(tmp_path, text, match): @@ -748,62 +736,19 @@ def test_radix_config_rejects(tmp_path, text, match): def test_radix_config_env_singleton_reloads_on_change(tmp_path, monkeypatch): - """`get_radixshmem_config` follows GLOBAL_CONFIG_FROM_ENV: the path and the - two per-node overrides; a test-installed config wins until reverted.""" + """`get_radixshmem_config` follows GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path; + a test-installed config wins until reverted.""" from flexkv.common.radixshmem_config import get_radixshmem_config set_radixshmem_config(None) monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", None) - monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_node_name", "") - monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_rpc_address", "") - assert get_radixshmem_config().cluster_id == "flexkv" - path = _write_yaml(tmp_path, "cluster:\n cluster_id: other\n") + assert get_radixshmem_config().server_name == "/flexkv" + path = _write_yaml(tmp_path, "server:\n name: /other\n") monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", path) - assert get_radixshmem_config().cluster_id == "other" - monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radix_node_name", "n7") - assert get_radixshmem_config().local_id == "other_n7" - set_radixshmem_config(_radix_config(cluster_id="pinned")) - assert get_radixshmem_config().cluster_id == "pinned" + assert get_radixshmem_config().server_name == "/other" + set_radixshmem_config(_radix_config(name="/pinned")) + assert get_radixshmem_config().server_name == "/pinned" set_radixshmem_config(None) - assert get_radixshmem_config().local_id == "other_n7" - - -def test_server_config_takes_the_yaml_sections(tmp_path): - """`build_radix_server_config` passes the four sections through and keeps - the geometry / naming its own.""" - model_config, cache_config = _configs(num_cpu_blocks=64) - cfg_path = _write_yaml(tmp_path, """ -cluster: - cluster_id: yamlsrv - expected_min_nodes: 2 - registry: etcd://10.0.0.1:2379 - rpc_interface: eth0 - index_dev: mlx5_3 - rht_transport: dc - num_rht_shards: 1 -data: - transfer_devices: [mlx5_4] - prefault: false - max_pending_jobs: 7 -index: - data_pool_ratio: 5.5 - register_chunk_size: 64 -server: - rpc_workers: 3 - endpoint: unix:///dev/shm/yamlsrv.sock -""") - rcfg = load_radixshmem_config(cfg_path) - cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) - assert cfg.index.name == "/shmradix_yamlsrv_cpu" and cfg.index.full_slots == 64 - assert cfg.index.data_pool_ratio == 5.5 - assert cfg.index.register_chunk_size == 64 # the file wins over the derived default - assert cfg.resolved_data_name == "/shmradix_yamlsrv_cpu_data" - assert cfg.data.transfer_devices == ["mlx5_4"] and cfg.data.max_pending_jobs == 7 - assert cfg.data.prefault is False and cfg.data.full_slot_bytes == 32768 - assert cfg.cluster.expected_min_nodes == 2 and cfg.cluster.index_dev == "mlx5_3" - assert cfg.cluster.rht_transport == "dc" and cfg.cluster.peer_index_transport == "xrc" - assert cfg.cluster.rht_slots_per_bucket == 4 and cfg.cluster.node_name == "" - assert cfg.rpc_workers == 3 and cfg.endpoint == "unix:///dev/shm/yamlsrv.sock" - assert cfg.distributed + assert get_radixshmem_config().te_server_id == "other" # ============================================================================= @@ -1317,7 +1262,7 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, from flexkv.common.config import CacheConfig, ModelConfig, SWAPoolConfig - rcfg = _radix_config(cluster_id=f"swaplanner{os.getpid()}") + rcfg = _radix_config(name=f"/swaplanner{os.getpid()}") saved = {"enable_radixshmem": GLOBAL_CONFIG_FROM_ENV.enable_radixshmem} GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True set_radixshmem_config(rcfg) @@ -1338,8 +1283,14 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, model_config = ModelConfig(num_layers=2, num_kv_heads=4, head_size=64, dtype=torch.float16, tp_size=1, dp_size=1) - cfg = bootstrap.build_radix_server_config(model_config, cache_config, rcfg) - _sweep_region(cfg.index.name, cfg.resolved_data_name) + # The operator's server: a budget sized so the planned counts are the + # test's; the planner's client brings the geometry. + geo = bootstrap.expected_geometry(model_config, cache_config) + data_bytes = num_blocks * geo.full_slot_bytes + swa_slots * geo.swa_slot_bytes + cfg = shmradix.ServerConfig(name=rcfg.server_name, data_bytes=data_bytes, + swa_ratio=swa_slots * geo.swa_slot_bytes / data_bytes, + prefault=False) + _sweep_region(cfg.name, cfg.resolved_data_name) server = shmradix.RadixServer(cfg).start() engine = RadixShmemCacheEngine(cache_config, model_config) assert engine.swa_op_constructor.enabled, \ @@ -1698,20 +1649,17 @@ def _node_main(rank, prefix, cluster_id, registry, rdma_dev, ready, done, output node_name=f"r{rank}", rpc_address="0.0.0.0", index_dev=rdma_dev, gid_idx=int(os.getenv("FLEXKV_TEST_RADIX_GID_IDX", "3")), bootstrap_timeout_sec=60, rht_slots_per_bucket=4) - cfg = shmradix.RadixServerConfig( - index=shmradix.IndexConfig(name=prefix, tokens_per_block=16, - full_slots=PEER_BLOCKS), - data=shmradix.DataPlaneConfig( - data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, full_slot_bytes=PEER_SLOT_BYTES, - slot_align=4096, data_name=data_name, prefault=False, - transfer_devices=[rdma_dev]), - cluster=shmradix.ClusterConfig(**cluster_kwargs), - endpoint=endpoint, + cfg = shmradix.ServerConfig( + name=prefix, data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, slot_align=4096, + data_name=data_name, prefault=False, transfer_devices=[rdma_dev], + endpoint=endpoint, cluster=shmradix.ClusterConfig(**cluster_kwargs), ) - server = shmradix.RadixServer(cfg).start() # collective: waits for both - set_radixshmem_config(_radix_config().replace_server(endpoint=endpoint)) - engine = CacheEngineRadixShmem(prefix, num_total_blocks=PEER_BLOCKS, - tokens_per_block=16, peer_enabled=True) + server = shmradix.RadixServer(cfg).start() # waiting: the engine's geometry starts the rendezvous + set_radixshmem_config(_radix_config(name=prefix, endpoint=endpoint, ready_timeout_s=180.0)) + engine = CacheEngineRadixShmem( + prefix, geometry=shmradix.Geometry(block_size=16, full_slot_bytes=PEER_SLOT_BYTES, + slot_align=4096), + num_total_blocks=PEER_BLOCKS, tokens_per_block=16, peer_enabled=True) if not engine.peer_enabled: raise RuntimeError("engine did not see a distributed region") cluster_rank = bootstrap.radix_cluster_rank(engine.client) From 69ca1af686065cf9d1f38fa27a3f995915404cff Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 13:11:37 +0800 Subject: [PATCH 11/21] radixshmem: the RHT registration chunk is the radix-server's --register-chunk-tokens FlexKV used to bring its own registration chunk: REGISTER_CHUNK_TOKENS = 4096 became index.register_chunk_size (in blocks) on the server config it built. radixshmem f939910 takes the chunk in tokens (Geometry.register_chunk_tokens, radix-server --register-chunk-tokens, default 4096) and publishes it with the geometry, so FlexKV now follows the server instead of dictating a value: - RadixGeometry.register_chunk_tokens (default 0 = the server's) is forwarded verbatim by to_shmradix(); expected_geometry leaves it at 0. - adopt_geometry also takes over the published register_chunk_tokens and its size in FlexKV blocks (radixshmem's rule, tokens // tokens_per_block and at least 1: register_chunk_blocks()); CacheEngineRadixShmem exposes both as attributes; the attach and adopt logs name them. - check_geometry compares a pinned value with the server's and warns when the server's chunk is not a whole number of FlexKV blocks. - tests: the adoption and engine tests run against a server started with a non-default chunk (2048 / 64 tokens) and cover the pinned and mismatch paths; docs/radixshmem/config_zh.md and CHANGELOG updated. --- CHANGELOG.md | 1 + docs/radixshmem/config_zh.md | 13 +++-- flexkv/cache/radix_shmem_engine.py | 8 ++- flexkv/server/shm_radix_bootstrap.py | 59 +++++++++++++++++---- tests/radixshmem/test_radix_shmem_engine.py | 52 ++++++++++++++++-- 5 files changed, 112 insertions(+), 21 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index cd03b5afc..065762f97 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 Universal: - radixshmem mode attaches to an operator-run `radix-server` and no longer creates one: the server is started per node with a name and a byte budget (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`), FlexKV's clients hand it the geometry (`shmradix.RadixClient(name, Geometry)`: tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), the server plans the slot counts from its budget and every FlexKV process adopts them into `CacheConfig.num_cpu_blocks` / `swa.num_slots` (`shm_radix_bootstrap.adopt_geometry`; `cpu_cache_gb` no longer sizes the CPU tier in this mode). `RadixServerProcess`, `build_radix_server_config`, `FLEXKV_RADIX_SERVER_LAUNCH_MODE`, `FLEXKV_RADIX_NODE_NAME` and `FLEXKV_RADIX_RPC_ADDRESS` are gone; the radixshmem YAML shrinks to `server` (`name`, `endpoint`, `ready_timeout_s`) and `client`, and rejects the former `cluster` / `data` / `index` sections (they are radix-server flags now; migration table in `docs/radixshmem/config_zh.md`). A server serving another geometry is refused (`GeometryMismatch`), so every engine on one server runs the same model, page size and SWA configuration. Requires radixshmem e5ce067 or later. +- radixshmem mode brings no RHT registration chunk of its own any more (the former `REGISTER_CHUNK_TOKENS = 4096` constant, handed to the server as `index.register_chunk_size` in blocks): the chunk is the radix-server's `--register-chunk-tokens` (radixshmem's default 4096 tokens). `RadixGeometry.register_chunk_tokens` forwards a pinned value (0 = the server's), `adopt_geometry` and `CacheEngineRadixShmem.register_chunk_tokens` / `register_chunk_blocks` take the published value over, converted to blocks by radixshmem's rule (`tokens // tokens_per_block`, at least 1); `check_geometry` compares a pinned value and warns when the server's chunk is not a whole number of FlexKV blocks. - radixshmem mode now uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. The radix-server runs as a subprocess of the bootstrap DP (`FLEXKV_RADIX_SERVER_LAUNCH_MODE`). radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). See `docs/radixshmem_integration.md` and `docs/radixshmem_cross_node.md` - radixshmem mode is configured by one YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `cluster` / `data` / `index` / `server` sections pass through by key to radixshmem's `ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig` (validated against the installed dataclasses, geometry keys rejected), a `client` section holds the prefetch limits; `index.register_chunk_size` defaults to `4096 / tokens_per_block` blocks (one RHT registration chunk per 4096 tokens). The file is global: `cluster.cluster_id` is the only namespace (etcd keys and every shm / socket / TE channel name), node identity derives from `cluster.rpc_interface`. The `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone; `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` remain as per-node overrides for co-located nodes. Examples in `examples/radixshmem_configs/`, reference `docs/radixshmem/config_zh.md`. - radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles now roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not. Peer reuse in this mode follows the radixshmem YAML (`distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 181da46db..202e35376 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -86,19 +86,21 @@ FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `sh | `full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 CPU block 字节数(每 PP 段层数 × 节点内 KV head 数 × head_size × kv_dim × dtype × tokens_per_block) | | `swa_slot_bytes` / `swa_window_blocks` | `CacheConfig.swa` 开启时:一个 SWA page 的字节数(uint8)与窗口块数;未开启则没有 SWA 池 | | `slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,保证 SlotStore stride 等于 block 字节数 | +| `register_chunk_tokens` | 不传(0):RHT 注册粒度由 server 的 `--register-chunk-tokens` 决定(radixshmem 默认 4096 token)。FlexKV 不再有自己的 4096 常量,attach 后采纳 server 发布的值,并按 radixshmem 的规则换算成 block 数(`tokens // tokens_per_block`,至少 1):`adopt_geometry` 返回的 `register_chunk_tokens` / `register_chunk_blocks`,`CacheEngineRadixShmem` 的同名属性。`RadixGeometry.register_chunk_tokens` 非 0 时原样交给 server(pin) | server 收到几何后的规划(radixshmem 的规则):`swa_slots = floor(swa_ratio × data_bytes / swa_stride)`, `full_slots = (data_bytes − SWA 占用) / full_stride`。任一池算出 0 个 slot、模型有 SWA 而 `--swa-ratio` 为 0, 都在 configure 时拒绝,FlexKV 报 `cannot serve FlexKV's geometry`。 **采纳**:attach 成功后 `adopt_geometry` 把 `pools.full.num_slots` 写进 `CacheConfig.num_cpu_blocks`, -`pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,日志形如 -`adopted radix-server /flexkv's slot counts: FULL 8605 slots (cpu_cache_gb had given 1524), SWA 1024 slots`。 +`pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,并带回 server 的 `register_chunk_tokens` 及其 block 数,日志形如 +`adopted radix-server /flexkv's geometry: FULL 8605 slots (cpu_cache_gb had given 1524), SWA 1024 slots (had 1024); RHT registration chunk 4096 tokens = 64 blocks`。 之后 TE 的 StorageEngine、cache engine、指标都用采纳后的值。 **校验**:每个 attach 方(KVManager、cache engine、TE)用 `check_geometry` 复核 server 发布的 `block_size`、各池 `slot_bytes`、SlotStore stride、SWA 窗口与自己的布局一致,不一致报错退出,不会静默错位传输。 -slot 数不在校验范围内,它们是 server 的。 +slot 数不在校验范围内,它们是 server 的;`register_chunk_tokens` 只在 FlexKV pin 了值时比对,server 的值不是 +`tokens_per_block` 的整数倍时只告警(chunk 取整到整 block)。 **同一 server 上的多个 client** 必须带相同的几何:相同模型、page size、SWA 配置。第二个不同的几何被 server 以 `GeometryMismatch` 拒绝,FlexKV 报 `already serves another geometry`。单机下 TP 不同的同一模型通常几何相同 @@ -136,7 +138,7 @@ radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 \ --transfer-dev mlx5_1 --transfer-dev mlx5_2 --bootstrap-timeout 600 ``` -集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`)由第一个拿到几何的节点发布到 +集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`、`register_chunk_tokens`)由第一个拿到几何的节点发布到 etcd `radix//geometry/`,其余 `waiting` 的节点采纳;各节点的 slot 数可以不同(预算可以不同)。 FlexKV 侧每个节点同一份 YAML 即可。`ready_timeout_s` 要不小于 `--bootstrap-timeout`。 @@ -196,7 +198,8 @@ server 一直在等几何或配置失败报 `not ready within ...`(带 server | `cluster.index_dev` / `gid_idx` / `rht_transport` / `peer_index_transport` / `remote_op_transport` / `zmq_listen_port` | 同名 `--index-dev` 等 | | `data.transfer_devices` / `transfer_protocol` / `transfer_ip` / `transfer_port` / `transfer_metadata` | `--transfer-dev`(可重复)/ `--transfer-protocol` / `--transfer-ip` / `--transfer-port` / `--transfer-metadata` | | `data.prefault` / `max_inflight` / `max_pending_jobs` / `job_ttl_s` | `--no-prefault` / `--max-inflight` / `--max-pending-jobs` / `--job-ttl` | -| `index.data_pool_ratio` / `background_evict_ratio` / `max_nodes` / `register_chunk_size` | `--data-pool-ratio` / `--background-evict-ratio` / `--max-nodes` / `--register-chunk-tokens`(按 token 数) | +| `index.data_pool_ratio` / `background_evict_ratio` / `max_nodes` | `--data-pool-ratio` / `--background-evict-ratio` / `--max-nodes` | +| `index.register_chunk_size`(block 数;FlexKV 曾固定按 4096 token 换算) | `--register-chunk-tokens`(token 数,默认 4096);FlexKV 采纳 server 的值,见第 3 节 | | `server.endpoint` | 保留:server 的 `--endpoint` 与 YAML `server.endpoint` 各写一次 | | `server.rpc_workers` / `hugepage_path` | `--rpc-workers` / `--hugepage-path` | | (由 FlexKV 推导的 slot 数、`data_bytes`) | slot 数由 `--data-bytes` 和 `--swa-ratio` 决定,FlexKV 采纳 | diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 1a0692142..37ac4fcd1 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -164,7 +164,7 @@ def __init__(self, the server is part of a cluster; False switches it off. `num_total_blocks` is FlexKV's expectation; the region's capacity is authoritative.""" - from flexkv.server.shm_radix_bootstrap import attach_radix_client + from flexkv.server.shm_radix_bootstrap import attach_radix_client, register_chunk_blocks self.event_collector = event_collector self._metrics_collector = metrics_collector @@ -190,6 +190,12 @@ def __init__(self, f"radix-server {self.shm_name} has tokens_per_block={region_tpb}, " f"FlexKV is configured with {tokens_per_block}") self.tokens_per_block = int(tokens_per_block) + # The RHT registration chunk is the server's (--register-chunk-tokens); + # FlexKV adopts it rather than bringing one of its own. + self.register_chunk_tokens = int( + (self._client.geometry or {}).get("register_chunk_tokens", 0)) + self.register_chunk_blocks = register_chunk_blocks(self.register_chunk_tokens, + self.tokens_per_block) capacity = int(self._client.mempool_total()) if num_total_blocks > 0 and capacity != int(num_total_blocks): flexkv_logger.warning( diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index e8e5758e1..d98fba10d 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -159,11 +159,17 @@ def slot_align_for(*sizes: int) -> int: @dataclasses.dataclass(frozen=True) class RadixGeometry: """FlexKV's side of the geometry: what one slot of each pool must hold. The - slot counts are not here; the server plans them from its byte budget.""" + slot counts are not here; the server plans them from its byte budget. Nor + is the RHT registration chunk unless pinned: ``register_chunk_tokens`` 0 + leaves it to the server's ``--register-chunk-tokens`` (radixshmem's default, + 4096 tokens) and FlexKV adopts the published value (:func:`adopt_geometry`, + :func:`register_chunk_blocks`).""" tokens_per_block: int full_slot_bytes: int swa_slot_bytes: int = 0 swa_window_blocks: int = 0 + # RHT registration granularity in tokens; 0 = the server's --register-chunk-tokens. + register_chunk_tokens: int = 0 @property def has_swa(self) -> bool: @@ -183,13 +189,27 @@ def to_shmradix(self) -> "shmradix.Geometry": swa_slot_bytes=int(self.swa_slot_bytes), swa_window_blocks=int(self.swa_window_blocks) if self.has_swa else 0, slot_align=int(self.slot_align), + register_chunk_tokens=int(self.register_chunk_tokens), ) def describe(self) -> str: s = f"tokens_per_block={self.tokens_per_block}, FULL slot {self.full_slot_bytes} B" if self.has_swa: s += f", SWA slot {self.swa_slot_bytes} B (window {self.swa_window_blocks})" - return s + f", slot_align={self.slot_align}" + s += f", slot_align={self.slot_align}" + if self.register_chunk_tokens: + s += f", register_chunk_tokens={self.register_chunk_tokens}" + return s + + +def register_chunk_blocks(register_chunk_tokens: int, tokens_per_block: int) -> int: + """The RHT registration chunk in blocks, by radixshmem's rule + (``ShmConfig::effective_register_chunk_size``): ``register_chunk_tokens // + block_size``, at least 1. 0 tokens = no chunk alignment.""" + tokens = int(register_chunk_tokens) + if tokens <= 0: + return 0 + return max(1, tokens // max(1, int(tokens_per_block))) def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> RadixGeometry: @@ -301,7 +321,8 @@ def attach_radix_client(name: Optional[str] = None, def _describe_published(g: Optional[Dict[str, Any]]) -> str: if not g: return "(none)" - parts = [f"block_size={g.get('block_size')}"] + parts = [f"block_size={g.get('block_size')}", + f"register_chunk_tokens={g.get('register_chunk_tokens', 0)}"] for kind, pool in (g.get("pools") or {}).items(): s = f"{kind.upper()} {pool.get('num_slots')} x {pool.get('slot_bytes')} B" if kind == "swa": @@ -334,6 +355,15 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, diffs: List[str] = [] if int(g["block_size"]) != expected.tokens_per_block: diffs.append(f"tokens_per_block server={g['block_size']} flexkv={expected.tokens_per_block}") + chunk_tokens = int(g.get("register_chunk_tokens", 0)) + if expected.register_chunk_tokens and chunk_tokens != expected.register_chunk_tokens: + diffs.append(f"register_chunk_tokens server={chunk_tokens} " + f"flexkv={expected.register_chunk_tokens}") + elif chunk_tokens % max(1, expected.tokens_per_block): + flexkv_logger.warning( + f"{label}: radix-server {client.name}'s register_chunk_tokens={chunk_tokens} is not a " + f"multiple of tokens_per_block={expected.tokens_per_block}; the RHT registration " + f"chunk is {register_chunk_blocks(chunk_tokens, expected.tokens_per_block)} blocks") full = pools["full"] if int(full["slot_bytes"]) != expected.full_slot_bytes: diffs.append(f"FULL slot_bytes server={full['slot_bytes']} flexkv={expected.full_slot_bytes}") @@ -371,12 +401,16 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", label: str = "radixshmem") -> Dict[str, int]: - """Take the slot counts the server planned from its budget over into - ``cache_config``: ``num_cpu_blocks`` = the FULL pool, ``swa.num_slots`` = - the SWA pool. Whatever ``cpu_cache_gb`` had produced was a placeholder in - this mode. Run it after any ``recompute_cache_block_counts`` in the same - process (that recompute sizes from ``cpu_cache_gb`` and would undo this). - Returns ``{"full": n, "swa": m}``.""" + """Take over what the server planned and published: the slot counts into + ``cache_config`` (``num_cpu_blocks`` = the FULL pool, ``swa.num_slots`` = + the SWA pool; whatever ``cpu_cache_gb`` had produced was a placeholder in + this mode) and the RHT registration chunk, which FlexKV does not bring + itself: ``register_chunk_tokens`` is the server's ``--register-chunk-tokens`` + and ``register_chunk_blocks`` that in FlexKV blocks. Run it after any + ``recompute_cache_block_counts`` in the same process (that recompute sizes + from ``cpu_cache_gb`` and would undo this). Returns ``{"full": n, "swa": m, + "register_chunk_tokens": t, "register_chunk_blocks": b}`` ("swa" only with + an SWA tier).""" g = _published_geometry(client, label) pools = g["pools"] counts: Dict[str, int] = {"full": int(pools["full"]["num_slots"])} @@ -394,7 +428,12 @@ def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", swa_before = int(swa.num_slots) swa.num_slots = counts["swa"] note += f", SWA {counts['swa']} slots (had {swa_before})" - flexkv_logger.info(f"{label}: adopted radix-server {client.name}'s slot counts: {note}") + counts["register_chunk_tokens"] = int(g.get("register_chunk_tokens", 0)) + counts["register_chunk_blocks"] = register_chunk_blocks(counts["register_chunk_tokens"], + int(g["block_size"])) + note += (f"; RHT registration chunk {counts['register_chunk_tokens']} tokens = " + f"{counts['register_chunk_blocks']} blocks") + flexkv_logger.info(f"{label}: adopted radix-server {client.name}'s geometry: {note}") return counts diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 5bbb5f3ef..4296b8b53 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -253,6 +253,20 @@ def test_take_insert_match_recycle(env): engine.recycle(free_slots) +def test_engine_adopts_the_servers_register_chunk(env): + """The RHT registration chunk is whatever the radix-server was started + with (--register-chunk-tokens); the engine carries it in tokens and in + FlexKV blocks, converted the way radixshmem does.""" + engine, _server = env.make("/cers_chunk") # radixshmem's default: 4096 tokens + assert engine.register_chunk_tokens == 4096 + assert engine.register_chunk_blocks == 4096 // 4 # tokens_per_block 4 + name = f"/cers_chunk{os.getpid()}" + cfg, geo = _server_config(name, blocks=64, tokens_per_block=4) + env.server(dataclasses.replace(cfg, register_chunk_tokens=64)) + engine2 = env.engine(name, geometry=geo, num_total_blocks=64, tokens_per_block=4) + assert engine2.register_chunk_tokens == 64 and engine2.register_chunk_blocks == 16 + + def test_insert_publishes_immediately(env): """There is no ready bit: being in the tree IS being servable. @@ -605,6 +619,25 @@ def test_expected_geometry_mirrors_the_storage_engine_layout(): assert set(spec["pools"]) == {"full"} +def test_register_chunk_is_the_servers_unless_pinned(): + """FlexKV brings no RHT registration chunk of its own: the geometry carries + 0 and the server's --register-chunk-tokens decides. A pinned value travels + verbatim. Tokens become blocks by radixshmem's rule: tokens // block_size, + at least 1, 0 = unaligned.""" + model_config, cache_config = _configs(num_cpu_blocks=64) + geo = bootstrap.expected_geometry(model_config, cache_config) + assert geo.register_chunk_tokens == 0 + assert geo.to_shmradix().to_dict()["register_chunk_tokens"] == 0 + assert "register_chunk" not in geo.describe() + pinned = dataclasses.replace(geo, register_chunk_tokens=2048) + assert pinned.to_shmradix().to_dict()["register_chunk_tokens"] == 2048 + assert "register_chunk_tokens=2048" in pinned.describe() + assert bootstrap.register_chunk_blocks(4096, 16) == 256 + assert bootstrap.register_chunk_blocks(4096, 4) == 1024 + assert bootstrap.register_chunk_blocks(100, 64) == 1 + assert bootstrap.register_chunk_blocks(0, 16) == 0 + + def test_client_brings_the_geometry_and_adopts_the_counts(env): """The operator's server knows only its budget; FlexKV's first client hands it the slot shape, the server plans the counts, `check_geometry` verifies @@ -616,6 +649,7 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): data_bytes = 64 * geo.full_slot_bytes + 16 * geo.swa_slot_bytes env.server(shmradix.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=16 * geo.swa_slot_bytes / data_bytes, + register_chunk_tokens=2048, # the operator's, not FlexKV's prefault=False)) # slot_align comes with the geometry client = bootstrap.attach_radix_client(name, geometry=geo, timeout_s=60) try: @@ -624,10 +658,17 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): assert pools["full"]["num_slots"] == 64 and pools["swa"]["num_slots"] == 16 assert int(client.store.pool(FULL).slot_bytes) == geo.full_slot_bytes # exact stride assert int(client.store.pool(_SWA).slot_bytes) == geo.swa_slot_bytes - # cpu_cache_gb's placeholders give way to the server's counts + # cpu_cache_gb's placeholders give way to the server's counts; the RHT + # registration chunk comes along: 2048 tokens = 128 blocks of 16 + assert client.geometry["register_chunk_tokens"] == 2048 cache_config.num_cpu_blocks, cache_config.swa.num_slots = 7, 3 - assert bootstrap.adopt_geometry(cache_config, client, "test") == {"full": 64, "swa": 16} + assert bootstrap.adopt_geometry(cache_config, client, "test") == { + "full": 64, "swa": 16, "register_chunk_tokens": 2048, "register_chunk_blocks": 128} assert cache_config.num_cpu_blocks == 64 and cache_config.swa.num_slots == 16 + # a chunk FlexKV pinned differently from the server's fails closed too + with pytest.raises(ValueError, match="register_chunk_tokens"): + bootstrap.check_geometry( + client, dataclasses.replace(geo, register_chunk_tokens=4096), "test") assert bootstrap.radix_cluster_rank(client) == 0 and client.info.world_size == 1 # the connectors' prefetch gate asks the server the same question assert bootstrap.radix_server_is_distributed(_radix_config(name=name), timeout_s=30) is False @@ -637,9 +678,10 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): bootstrap.check_geometry( client, bootstrap.expected_geometry(model_config, cache_config), "test") # ...and a second client bringing another geometry is refused by the server - other = dataclasses.replace(geo, tokens_per_block=32) - with pytest.raises(ValueError, match="another geometry"): - bootstrap.attach_radix_client(name, geometry=other, timeout_s=60) + for other in (dataclasses.replace(geo, tokens_per_block=32), + dataclasses.replace(geo, register_chunk_tokens=4096)): + with pytest.raises(ValueError, match="another geometry"): + bootstrap.attach_radix_client(name, geometry=other, timeout_s=60) finally: client.close() From aa26bb5b59f292976a77c2bb471350415e8a64e0 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 14:08:42 +0800 Subject: [PATCH 12/21] radixshmem: harden the shm TE ring, PUT planning and the numpy hash path; consistent changelog Review follow-ups on feat/radixshmem (base main 738ddc1): - shm_channel: LAYERWISE has a wire index. _TT_NAMES lacked it, so the first layerwise completion raised KeyError inside the TE's result thread and every rank on the node stopped receiving completions. An unknown name now goes out as None with one warning; a submit record that does not unpickle is dropped with an error instead of being re-read on every poll; a graph the TE cannot submit is reported back to its channel as failed (op_id -1) and a delivery error on one channel no longer stops the others. - radixshmem PUT: CacheEngineRadixShmem.take clamps a request to the pool's size (radixshmem's allocate_slots raises ValueError above it) and answers empty on a refusal; _plan_put returns the taken FULL/SWA slots and releases the match pin when planning fails after the take, then re-raises. Before, a server planning fewer SWA slots than the window leaked FULL slots on every PUT and left the matched prefix pinned for good. - gen_hashes_numpy checks its buffers (int64 tokens, uint64 hashes, both C-contiguous, enough tokens for the requested blocks) and gen_hashes raises TypeError on non-int64 input as the torch path did, instead of reading past an int32 buffer. Hashes of int64 input are unchanged. - KVManager.shutdown works on an instance whose __init__ did not reach the radixshmem setup (_shutdown_radix_shmem_children uses getattr); the CI unit file tests/test_kvmanager_client_api.py failed with AttributeError. - CHANGELOG: the radixshmem entries describe the final state only (operator-run server, server/client YAML, peer reuse from the server's world_size, the shm TE ring, numpy hashing). The stale bullets about the embedded server, the cluster/data/index YAML pass-through, `distributed` and two never-written docs are gone; the requirement is radixshmem f939910 (e5ce067 was amended away and is unreachable). Tests: tests/test_shm_channel.py (every TransferType round-trips, unknown type, poisoned record), tests/radixshmem/test_radix_shmem_engine.py (take clamps, PUT planning failure rolls back), new unit-scope tests/test_hash_utils.py (chain equivalence with Hasher, dtype and buffer checks). 472 passed across the CI unit files and the radixshmem suites. --- CHANGELOG.md | 10 +- csrc/bindings.cpp | 24 +++ flexkv/cache/radix_shmem_engine.py | 32 +++- flexkv/cache/radix_shmem_planner.py | 160 +++++++++++--------- flexkv/common/hash_utils.py | 7 +- flexkv/kvmanager.py | 7 +- flexkv/transfer/shm_channel.py | 47 +++++- flexkv/transfer/shm_channel_handle.py | 32 +++- tests/radixshmem/test_radix_shmem_engine.py | 49 ++++++ tests/test_hash_utils.py | 51 +++++++ tests/test_shm_channel.py | 54 +++++++ 11 files changed, 385 insertions(+), 88 deletions(-) create mode 100644 tests/test_hash_utils.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 065762f97..14731053f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,12 +10,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Feature Universal: -- radixshmem mode attaches to an operator-run `radix-server` and no longer creates one: the server is started per node with a name and a byte budget (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`), FlexKV's clients hand it the geometry (`shmradix.RadixClient(name, Geometry)`: tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), the server plans the slot counts from its budget and every FlexKV process adopts them into `CacheConfig.num_cpu_blocks` / `swa.num_slots` (`shm_radix_bootstrap.adopt_geometry`; `cpu_cache_gb` no longer sizes the CPU tier in this mode). `RadixServerProcess`, `build_radix_server_config`, `FLEXKV_RADIX_SERVER_LAUNCH_MODE`, `FLEXKV_RADIX_NODE_NAME` and `FLEXKV_RADIX_RPC_ADDRESS` are gone; the radixshmem YAML shrinks to `server` (`name`, `endpoint`, `ready_timeout_s`) and `client`, and rejects the former `cluster` / `data` / `index` sections (they are radix-server flags now; migration table in `docs/radixshmem/config_zh.md`). A server serving another geometry is refused (`GeometryMismatch`), so every engine on one server runs the same model, page size and SWA configuration. Requires radixshmem e5ce067 or later. +- radixshmem mode: the CPU tier is an operator-run `radix-server` (the radixshmem project; requires radixshmem f939910 or later). FlexKV no longer creates a server: the operator starts one per node with a name and a byte budget (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`), FlexKV's clients hand it the geometry (`shmradix.RadixClient(name, Geometry)`: tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), the server plans the slot counts from its budget and every FlexKV process adopts them into `CacheConfig.num_cpu_blocks` / `swa.num_slots` (`shm_radix_bootstrap.adopt_geometry`; `cpu_cache_gb` no longer sizes the CPU tier in this mode). A server serving another geometry is refused (`GeometryMismatch`), so every engine on one server runs the same model, page size and SWA configuration. `RadixServerProcess`, `build_radix_server_config`, `FLEXKV_RADIX_SERVER_LAUNCH_MODE`, `FLEXKV_RADIX_NODE_NAME`, `FLEXKV_RADIX_RPC_ADDRESS`, the `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone. +- radixshmem mode is configured by one small YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `server` (`name`, `endpoint`, `ready_timeout_s`: which radix-server to attach to and how long to wait for it to be reachable and ready) and `client` (`prefetch_timeout_ms`, `prefetch_max_inflight`, `max_outstanding`). No file means `radix-server --name /flexkv` on the local socket. The former `cluster` / `data` / `index` sections are rejected: those settings are radix-server command-line flags now (migration table in `docs/radixshmem/config_zh.md`, launch scripts and a YAML in `examples/radixshmem_configs/`). FlexKV's own TE channel names derive from `server.name`. - radixshmem mode brings no RHT registration chunk of its own any more (the former `REGISTER_CHUNK_TOKENS = 4096` constant, handed to the server as `index.register_chunk_size` in blocks): the chunk is the radix-server's `--register-chunk-tokens` (radixshmem's default 4096 tokens). `RadixGeometry.register_chunk_tokens` forwards a pinned value (0 = the server's), `adopt_geometry` and `CacheEngineRadixShmem.register_chunk_tokens` / `register_chunk_blocks` take the published value over, converted to blocks by radixshmem's rule (`tokens // tokens_per_block`, at least 1); `check_geometry` compares a pinned value and warns when the server's chunk is not a whole number of FlexKV blocks. -- radixshmem mode now uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. The radix-server runs as a subprocess of the bootstrap DP (`FLEXKV_RADIX_SERVER_LAUNCH_MODE`). radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). See `docs/radixshmem_integration.md` and `docs/radixshmem_cross_node.md` -- radixshmem mode is configured by one YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `cluster` / `data` / `index` / `server` sections pass through by key to radixshmem's `ClusterConfig` / `DataPlaneConfig` / `IndexConfig` / `RadixServerConfig` (validated against the installed dataclasses, geometry keys rejected), a `client` section holds the prefetch limits; `index.register_chunk_size` defaults to `4096 / tokens_per_block` blocks (one RHT registration chunk per 4096 tokens). The file is global: `cluster.cluster_id` is the only namespace (etcd keys and every shm / socket / TE channel name), node identity derives from `cluster.rpc_interface`. The `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone; `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` remain as per-node overrides for co-located nodes. Examples in `examples/radixshmem_configs/`, reference `docs/radixshmem/config_zh.md`. -- radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles now roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not. Peer reuse in this mode follows the radixshmem YAML (`distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. +- radixshmem mode uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. Peer reuse follows the radix-server itself (`world_size > 1`, asked through `shm_radix_bootstrap.radix_server_is_distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). Reference: `docs/radixshmem/config_zh.md`. +- radixshmem mode runs one node-local shm TE process per node (`TransferManagerShmTEProcess`, `te_shm_main`) that every DP rank, and with `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID` every co-located engine, reaches over shared-memory submit / result rings (`flexkv/transfer/shm_channel.py`; the result record carries the worker-measured `wait_ms` / `xfer_ms` / `e2e_ms`). Every `TransferType` value, `LAYERWISE` included, has a wire index and an unknown one degrades to None; a submit record the TE cannot decode is dropped with an error, a graph it cannot submit is reported back as failed, and a delivery error on one channel does not stop the others. Op and graph ids come from per-process ranges (`set_op_id_range` / `set_graph_id_range`, `dp_client_id << 32`), lock-free. +- radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not, and a PUT whose planning fails after its slots were taken returns them and drops the pin before re-raising. `CacheEngineRadixShmem.take` clamps a request to the pool's size instead of letting radixshmem refuse it. - `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`) is trimmed to what `RadixShmemCacheEngine` uses: the `CacheEngineAccel`-compatibility parameters (`device_type`, `evict_ratio`, `evict_start_threshold`, `hit_reward_seconds`, `eviction_policy`, `protected_threshold`, `tokens_per_block=-1`), `take(strict=)`, `match(gpu_matched_blocks=)`, the `mempool` view, `start()`, `store` / `cluster_rank` and the `FLEXKV_TRACE_RADIX_PEER` variable are gone (prefetch logs at debug level; the planner reports mempool metrics itself). +- `gen_hashes` and `Hasher.update` hash numpy buffers directly (`c_ext.gen_hashes_numpy` / `update_numpy`) instead of going through `torch.from_numpy`, which is not safe to call concurrently. Hashes are unchanged for int64 tokens; `gen_hashes` keeps requiring int64 and the binding now checks dtype, contiguity and sizes instead of reinterpreting the buffer. Targeting SGLang: - The native FlexKV backend is available in upstream SGLang `v0.5.16` and later; no patch is required ([sglang#29701](https://github.com/sgl-project/sglang/pull/29701)) diff --git a/csrc/bindings.cpp b/csrc/bindings.cpp index 64f94101f..e254eb7d5 100644 --- a/csrc/bindings.cpp +++ b/csrc/bindings.cpp @@ -487,8 +487,32 @@ PYBIND11_MODULE(c_ext, m) { [](flexkv::Hasher &hasher, py::array token_ids, int tokens_per_block, py::array block_hashes) { // numpy-buffer variant of gen_hashes; bypasses torch.from_numpy. + // Same contract as gen_hashes(Tensor): int64 tokens, uint64 hashes, + // both C-contiguous, one hash per whole block. Checked here because a + // raw buffer reinterpreted as int64 would otherwise be read past its + // end (an int32 array is half as long as the loop assumes). + if (tokens_per_block <= 0) + throw py::value_error("gen_hashes_numpy: tokens_per_block must be > 0"); + if (!py::isinstance>(token_ids)) + throw py::type_error( + "gen_hashes_numpy: token_ids must be an int64 array, got dtype " + + py::str(token_ids.dtype()).cast()); + if (!py::isinstance>(block_hashes)) + throw py::type_error( + "gen_hashes_numpy: block_hashes must be a uint64 array, got dtype " + + py::str(block_hashes.dtype()).cast()); + if (!(token_ids.flags() & py::array::c_style) || + !(block_hashes.flags() & py::array::c_style)) + throw py::value_error( + "gen_hashes_numpy: token_ids and block_hashes must be C-contiguous"); py::buffer_info tok = token_ids.request(); py::buffer_info bh = block_hashes.request(true); + if (bh.size * static_cast(tokens_per_block) > tok.size) + throw py::value_error( + "gen_hashes_numpy: block_hashes has " + std::to_string(bh.size) + + " blocks of " + std::to_string(tokens_per_block) + + " tokens but token_ids holds only " + std::to_string(tok.size) + + " tokens"); const int64_t *tok_ptr = static_cast(tok.ptr); flexkv::HashType *bh_ptr = static_cast(bh.ptr); for (py::ssize_t i = 0; i < bh.size; i++) { diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 37ac4fcd1..8019053b2 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -358,15 +358,39 @@ def take(self, component: shmradix.ComponentType = COMPONENT_FULL) -> np.ndarray: """Allocate up to `num_required_blocks` slots, evicting unpinned LRU blocks as needed; fewer come back when the pool cannot supply them - (the SWA pool is all-or-none).""" - slots = np.asarray(self._tree.allocate_slots(num_required_blocks, - component=component), - dtype=np.int64) + (the SWA pool is all-or-none). A request above the pool's size is + clamped to it: radixshmem refuses such a request outright + (`allocate_slots` raises), and a refusal here would leak whatever the + caller took before.""" + total = self._pool_total(component) + n = int(num_required_blocks) + if total is not None and n > total: + n = total + if n <= 0: + return _empty_i64() + try: + slots = np.asarray(self._tree.allocate_slots(n, component=component), + dtype=np.int64) + except ValueError as e: + # Refused by the index (component not enabled, pool gone): nothing + # was allocated, so an empty answer is the truthful one. + flexkv_logger.warning( + f"radix-server {self.shm_name} refused a {n}-slot {component} " + f"allocation: {e}") + return _empty_i64() if (self._metrics_collector is not None and len(slots) > 0 and component == COMPONENT_FULL): # SWA has its own pool self._metrics_collector.record_allocation("cpu", len(slots)) return slots + def _pool_total(self, component: shmradix.ComponentType) -> Optional[int]: + """Slots in `component`'s pool, None when the index has no such pool.""" + if component == COMPONENT_FULL: + return int(self._tree.mempool_total()) + if component == COMPONENT_SWA: + return int(self._tree.swa_mempool_total()) + return None + def recycle(self, physical_blocks: np.ndarray, component: shmradix.ComponentType = COMPONENT_FULL) -> None: diff --git a/flexkv/cache/radix_shmem_planner.py b/flexkv/cache/radix_shmem_planner.py index ebcfd5bce..5f95852b5 100644 --- a/flexkv/cache/radix_shmem_planner.py +++ b/flexkv/cache/radix_shmem_planner.py @@ -510,73 +510,97 @@ def _release_match() -> RadixPutPlan: return _release_match() swa_new: Optional[np.ndarray] = None - if self.swa_op_constructor.enabled: - k = min(block_mask_end, self.cache_config.swa.window_blocks) - swa_take = cpu_engine.take(num_required_blocks=k, component=COMPONENT_SWA) - if len(swa_take) == k: - swa_new = swa_take - else: - # All-or-none contract says this is empty; recycle defensively - # in case it ever is not. - cpu_engine.recycle(swa_take, component=COMPONENT_SWA) - flexkv_logger.warning( - f"radixshmem PUT {request_id}: no {k}-slot SWA window " - f"available; storing Full KV only" - ) - - transfer_graph = TransferOpGraph() - finished_ops_ids: List[int] = [] - - fragment_gpu_blocks = gpu_block_ids[num_skipped:] - op_d2h = TransferOp( - graph_id=transfer_graph.graph_id, - transfer_type=TransferType.D2H, - src_block_ids=fragment_gpu_blocks, - dst_block_ids=cpu_new, - dp_client_id=dp_client_id, - ) - transfer_graph.add_transfer_op(op_d2h) - finished_ops_ids.append(op_d2h.op_id) - - if swa_new is not None: - swa_ops = self.swa_op_constructor.build_put_chain( - transfer_graph, - gpu_slot_ids=np.zeros(len(swa_new), dtype=np.int64), - cpu_slot_ids=swa_new, + try: + if self.swa_op_constructor.enabled: + k = min(block_mask_end, self.cache_config.swa.window_blocks) + swa_take = cpu_engine.take(num_required_blocks=k, component=COMPONENT_SWA) + if len(swa_take) == k: + swa_new = swa_take + else: + # All-or-none contract says this is empty; recycle defensively + # in case it ever is not. + cpu_engine.recycle(swa_take, component=COMPONENT_SWA) + flexkv_logger.warning( + f"radixshmem PUT {request_id}: no {k}-slot SWA window " + f"available; storing Full KV only" + ) + + transfer_graph = TransferOpGraph() + finished_ops_ids: List[int] = [] + + fragment_gpu_blocks = gpu_block_ids[num_skipped:] + op_d2h = TransferOp( + graph_id=transfer_graph.graph_id, + transfer_type=TransferType.D2H, + src_block_ids=fragment_gpu_blocks, + dst_block_ids=cpu_new, dp_client_id=dp_client_id, - return_op_ids=True, ) - assert swa_ops.d2h_id is not None - finished_ops_ids.append(swa_ops.d2h_id) - - on_complete: List[Action] = [] - on_abort: List[Action] = [] - - def _arm(slots: np.ndarray, hold: Optional[Action], label: str, - component=COMPONENT_FULL) -> None: - staged = StagedRadixInsert(engine=cpu_engine, - sequence_meta=sequence_meta, - slots=slots, - path_end=block_mask_end, - label=label, - holds=[] if hold is None else [hold], - component=component) - on_complete.append(staged.publish) - on_abort.append(staged.abort) - - # The match pin travels with the LAST publish: SWA's insert refuses paths - # the Full tree does not reach yet, so it runs after FULL and releases. - _arm(cpu_new, cpu_match.release if swa_new is None else None, - f"PUT {request_id} CPU") - if swa_new is not None: - _arm(swa_new, cpu_match.release, f"PUT {request_id} CPU SWA", - component=COMPONENT_SWA) - - return RadixPutPlan( - transfer_graph=transfer_graph, - finished_ops_ids=finished_ops_ids, - num_gpu_blocks_to_transfer=len(fragment_gpu_blocks), - skipped_gpu_blocks=num_skipped, - on_complete=on_complete, - on_abort=on_abort, - ) + transfer_graph.add_transfer_op(op_d2h) + finished_ops_ids.append(op_d2h.op_id) + + if swa_new is not None: + swa_ops = self.swa_op_constructor.build_put_chain( + transfer_graph, + gpu_slot_ids=np.zeros(len(swa_new), dtype=np.int64), + cpu_slot_ids=swa_new, + dp_client_id=dp_client_id, + return_op_ids=True, + ) + assert swa_ops.d2h_id is not None + finished_ops_ids.append(swa_ops.d2h_id) + + on_complete: List[Action] = [] + on_abort: List[Action] = [] + + def _arm(slots: np.ndarray, hold: Optional[Action], label: str, + component=COMPONENT_FULL) -> None: + staged = StagedRadixInsert(engine=cpu_engine, + sequence_meta=sequence_meta, + slots=slots, + path_end=block_mask_end, + label=label, + holds=[] if hold is None else [hold], + component=component) + on_complete.append(staged.publish) + on_abort.append(staged.abort) + + # The match pin travels with the LAST publish: SWA's insert refuses paths + # the Full tree does not reach yet, so it runs after FULL and releases. + _arm(cpu_new, cpu_match.release if swa_new is None else None, + f"PUT {request_id} CPU") + if swa_new is not None: + _arm(swa_new, cpu_match.release, f"PUT {request_id} CPU SWA", + component=COMPONENT_SWA) + + return RadixPutPlan( + transfer_graph=transfer_graph, + finished_ops_ids=finished_ops_ids, + num_gpu_blocks_to_transfer=len(fragment_gpu_blocks), + skipped_gpu_blocks=num_skipped, + on_complete=on_complete, + on_abort=on_abort, + ) + except BaseException: + # Planning failed after slots were taken and before any handle could + # own them: hand them back, drop the match pin, then re-raise. (A + # StagedRadixInsert armed above is garbage now; nothing calls it.) + for slots, comp in ((cpu_new, COMPONENT_FULL), (swa_new, COMPONENT_SWA)): + if slots is None or len(slots) == 0: + continue + try: + if comp == COMPONENT_FULL: + cpu_engine.recycle(slots) + else: + cpu_engine.recycle(slots, component=comp) + except Exception as e: # noqa: BLE001 - report, keep unwinding + flexkv_logger.error( + f"radixshmem PUT {request_id}: could not return {len(slots)} " + f"{comp} slots after a planning failure: {e!r}") + try: + cpu_match.release() + except Exception as e: # noqa: BLE001 + flexkv_logger.error( + f"radixshmem PUT {request_id}: could not release the match pin " + f"after a planning failure: {e!r}") + raise diff --git a/flexkv/common/hash_utils.py b/flexkv/common/hash_utils.py index 1f0df25bc..784205061 100644 --- a/flexkv/common/hash_utils.py +++ b/flexkv/common/hash_utils.py @@ -41,10 +41,15 @@ def hash_array_with_prefix(array: np.ndarray, prefix: int) -> HashType: return HashType(_HASHER.digest()) def gen_hashes(token_ids: np.ndarray, tokens_per_block: int, hasher: Optional[Hasher] = None) -> np.ndarray: + token_ids = np.ascontiguousarray(token_ids) + if token_ids.dtype != np.int64: + # The hash covers the 8-byte token words; another width would be + # reinterpreted, not converted. Same contract as the torch path had. + raise TypeError(f"gen_hashes: token_ids must be int64, got {token_ids.dtype}") block_hashes = np.zeros(token_ids.size // tokens_per_block, dtype=np.uint64) if hasher is None: hasher = Hasher() - c_ext.gen_hashes_numpy(hasher.hasher, np.ascontiguousarray(token_ids), tokens_per_block, block_hashes) + c_ext.gen_hashes_numpy(hasher.hasher, token_ids, tokens_per_block, block_hashes) return block_hashes if __name__ == "__main__": diff --git a/flexkv/kvmanager.py b/flexkv/kvmanager.py index 77444d909..838043725 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -253,9 +253,12 @@ def _init_radix_shmem_path(self, raise def _shutdown_radix_shmem_children(self) -> None: - if self._shm_te_process is not None: - self._shm_te_process.shutdown() + # getattr: shutdown() must work on a KVManager whose __init__ did not + # get this far (or was bypassed, as the unit tests do). + te = getattr(self, "_shm_te_process", None) + if te is not None: self._shm_te_process = None + te.shutdown() def _spawn_shm_te(self) -> None: """Bootstrap proc (local dp 0) only: spawn the TE subprocess every DP diff --git a/flexkv/transfer/shm_channel.py b/flexkv/transfer/shm_channel.py index e1196c771..efa12170c 100644 --- a/flexkv/transfer/shm_channel.py +++ b/flexkv/transfer/shm_channel.py @@ -130,20 +130,44 @@ def _round_up(x: int, m: int) -> int: _COMPLETED_OP = struct.Struct(" int: + """Wire index of a CompletedOp.transfer_type (a TransferType value, the + member itself, or None). A name the table does not know goes out as None + with one warning per name: the CE loses that op's type attribution, which + beats a KeyError in the TE's result thread that would stop every + completion on the node.""" + if tt is None: + return _TT_NONE + key = getattr(tt, "value", tt) + idx = _TT_NAME_TO_IDX.get(key) + if idx is None: + if key not in _tt_unknown_warned: + _tt_unknown_warned.add(key) + from flexkv.common.debug import flexkv_logger + flexkv_logger.warning( + f"shm channel: transfer type {key!r} has no wire index; it is sent as " + f"None (add it to shm_channel._TT_NAMES)") + return _TT_NONE + return idx def encode_completed_op(op: Any) -> bytes: """Pack a CompletedOp into its fixed-width record.""" - tt = op.transfer_type - tt_idx = _TT_NONE if tt is None else _TT_NAME_TO_IDX[tt] + tt_idx = transfer_type_index(op.transfer_type) flags = 1 if getattr(op, "failed", False) else 0 return _COMPLETED_OP.pack( op.graph_id, op.op_id, tt_idx, op.num_blocks, op.num_bytes, flags, @@ -158,7 +182,7 @@ def decode_completed_op(buf: Any, off: int) -> Any: from flexkv.common.transfer import CompletedOp graph_id, op_id, tt_idx, num_blocks, num_bytes, flags, wait_ms, xfer_ms, e2e_ms = \ _COMPLETED_OP.unpack_from(buf, off) - tt = None if tt_idx == _TT_NONE else _TT_NAMES[tt_idx] + tt = None if tt_idx == _TT_NONE or tt_idx >= len(_TT_NAMES) else _TT_NAMES[tt_idx] return CompletedOp( graph_id=graph_id, op_id=op_id, @@ -415,9 +439,18 @@ def submit_recv(self) -> List[Any]: off + _FRAG_HDR_SIZE + n])) rp = (rp + 1) & (slots - 1) if is_last: - out.append(pickle.loads(b"".join(frags))) + payload = b"".join(frags) frags.clear() - self._submit_r.value = rp # release this message's slots + self._submit_r.value = rp # release this message's slots (payload is a copy) + try: + out.append(pickle.loads(payload)) + except Exception as e: # noqa: BLE001 - a record we cannot decode + # Dropping it costs the sender one task; keeping it would + # re-raise on every poll and stall the whole ring. + from flexkv.common.debug import flexkv_logger + flexkv_logger.error( + f"shm channel {self.channel_id}: dropping an undecodable " + f"{len(payload)}-byte submit record ({e!r})") return out def result_send(self, ops: List[Any]) -> None: diff --git a/flexkv/transfer/shm_channel_handle.py b/flexkv/transfer/shm_channel_handle.py index 4b514475a..57af4c17b 100644 --- a/flexkv/transfer/shm_channel_handle.py +++ b/flexkv/transfer/shm_channel_handle.py @@ -235,7 +235,13 @@ def _poll_submits(self) -> None: graph = m.graph with self._owner_lock: self._graph_owner[graph.graph_id] = ch.channel_id - self._tm.submit(graph) + try: + self._tm.submit(graph) + except Exception as e: # noqa: BLE001 - one bad graph must not stop the node + flexkv_logger.error( + f"TE could not submit graph {graph.graph_id} from channel " + f"{ch.channel_id}: {e!r}; failing it", exc_info=True) + self._fail_graph(ch, graph.graph_id) if had_work: idle_spins = 0 continue @@ -282,8 +288,30 @@ def _poll_results(self) -> None: continue by_channel.setdefault(owner, []).append(op) for ch_id, ops in by_channel.items(): - if 0 <= ch_id < len(self._channels): + if not (0 <= ch_id < len(self._channels)): + continue + try: self._channels[ch_id].result_send(ops) + except Exception as e: # noqa: BLE001 - keep serving the other channels + flexkv_logger.error( + f"TE could not deliver {len(ops)} completion(s) to channel " + f"{ch_id} (graphs {sorted({op.graph_id for op in ops})}): {e!r}", + exc_info=True) + + def _fail_graph(self, ch, graph_id: int) -> None: + """Tell the owning CE that `graph_id` is over and failed when the TE could + not run it, so its task errors out instead of waiting forever.""" + from flexkv.common.transfer import CompletedOp + with self._owner_lock: + self._graph_owner.pop(graph_id, None) + try: + ch.result_send([CompletedOp(graph_id=graph_id, op_id=-1, transfer_type=None, + num_blocks=0, num_bytes=0, wait_ms=0.0, + xfer_ms=0.0, e2e_ms=0.0, failed=True)]) + except Exception as e: # noqa: BLE001 + flexkv_logger.error( + f"TE could not report the failure of graph {graph_id} to channel " + f"{ch.channel_id}: {e!r}") def te_shm_main(model_config: ModelConfig, diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 4296b8b53..ccaf904e7 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -267,6 +267,24 @@ def test_engine_adopts_the_servers_register_chunk(env): assert engine2.register_chunk_tokens == 64 and engine2.register_chunk_blocks == 16 +def test_take_clamps_to_the_pool(env): + """radixshmem refuses an allocation larger than the pool (allocate_slots + raises ValueError); the engine clamps instead, so a planner asking for more + than exists gets what exists and nothing it took before leaks.""" + engine, _server = env.make(f"/cers_clamp{os.getpid()}", blocks=64) + slots = engine.take(num_required_blocks=64 + 7) + assert 0 < len(slots) <= 64 + engine.recycle(slots) + assert engine.take(num_required_blocks=0).size == 0 + # the SWA pool is all-or-none: an oversize window comes back short and the + # planner stores Full KV only + swa_engine, _ = _make_swa_engine(env, f"/cers_clamp_swa{os.getpid()}", blocks=64, + swa_slots=SWA_W, window_blocks=SWA_W) + got = swa_engine.take(num_required_blocks=SWA_W + 3, component=_SWA) + assert len(got) <= SWA_W + swa_engine.recycle(got, component=_SWA) + + def test_insert_publishes_immediately(env): """There is no ready bit: being in the tree IS being servable. @@ -1245,6 +1263,37 @@ def test_get_abort_drops_the_pin(): assert engine.inserted_pools == [] # type: ignore[attr-defined] +def test_put_planning_failure_returns_the_taken_slots_and_drops_the_pin(): + """An exception after the FULL slots were taken (here the SWA take blows up) + must not leak them or the match pin: the planner recycles, releases and + re-raises. Before this, a server planning fewer SWA slots than the window + leaked FULL slots on every PUT.""" + engine = _global_cache_engine() + released = [] + _force_radixshmem(engine, _local_match(np.arange(20, 22), + finalize=lambda: released.append(1))) + tier = engine.cpu_cache_engine + free_before = tier.mempool.num_free_blocks + real_take = tier.take + + def _take(num_required_blocks, component=None, **kwargs): + if component is not None: + raise RuntimeError("SWA pool exploded") + return real_take(num_required_blocks=num_required_blocks, **kwargs) + + tier.take = _take # type: ignore[method-assign] + engine.swa_op_constructor = SimpleNamespace(enabled=True) # type: ignore[assignment] + engine.cache_config.swa = SimpleNamespace(window_blocks=2) + token_ids, token_mask, slot_mapping = _fake_request(5) + with pytest.raises(RuntimeError, match="SWA pool exploded"): + engine.put(request_id=3, token_ids=token_ids, token_mask=token_mask, + slot_mapping=slot_mapping, dp_client_id=0) + assert tier.mempool.num_free_blocks == free_before # the 3 FULL slots came back + assert [len(s) for s in engine.aborted_slots] == [3] # type: ignore[attr-defined] + assert released == [1] # pin dropped + assert engine.inserted_pools == [] # type: ignore[attr-defined] + + def test_put_abort_returns_the_staged_slots_and_drops_the_pin(): """A cancelled PUT never runs its D2H: nothing is published, the staged slots go back to the mempool and the match pin is released.""" diff --git a/tests/test_hash_utils.py b/tests/test_hash_utils.py new file mode 100644 index 000000000..b1499310a --- /dev/null +++ b/tests/test_hash_utils.py @@ -0,0 +1,51 @@ +"""flexkv.common.hash_utils: the numpy hashing path keeps the torch path's contract.""" +import numpy as np +import pytest + +from flexkv.common.hash_utils import Hasher, gen_hashes + +pytestmark = pytest.mark.unit + +TPB = 16 + + +def test_gen_hashes_chains_blocks_like_the_incremental_hasher(): + tokens = np.arange(1, 4 * TPB + 1, dtype=np.int64) + hashes = gen_hashes(tokens, TPB) + assert hashes.dtype == np.uint64 and hashes.shape == (4,) + h = Hasher() + expected = [] + for b in range(4): + h.update(tokens[b * TPB:(b + 1) * TPB]) + expected.append(int(h.digest())) + assert hashes.tolist() == expected + # a strided view hashes the same tokens as its contiguous copy + assert gen_hashes(np.repeat(tokens, 2)[::2], TPB).tolist() == expected + # a trailing partial block is not hashed + assert gen_hashes(tokens[:2 * TPB + 5], TPB).tolist() == expected[:2] + assert gen_hashes(tokens[:TPB - 1], TPB).size == 0 + + +def test_gen_hashes_rejects_non_int64_tokens(): + """int32 tokens reinterpreted as int64 would hash the neighbouring memory + (the torch path raised on them too).""" + with pytest.raises(TypeError, match="int64"): + gen_hashes(np.arange(2 * TPB, dtype=np.int32), TPB) + + +def test_gen_hashes_numpy_binding_checks_its_buffers(): + from flexkv import c_ext + tokens = np.arange(2 * TPB, dtype=np.int64) + out = np.zeros(2, dtype=np.uint64) + with pytest.raises(TypeError, match="int64"): + c_ext.gen_hashes_numpy(Hasher().hasher, tokens.astype(np.int32), TPB, out) + with pytest.raises(TypeError, match="uint64"): + c_ext.gen_hashes_numpy(Hasher().hasher, tokens, TPB, np.zeros(2, dtype=np.int64)) + with pytest.raises(ValueError, match="blocks"): # 3 blocks asked of 32 tokens + c_ext.gen_hashes_numpy(Hasher().hasher, tokens, TPB, np.zeros(3, dtype=np.uint64)) + with pytest.raises(ValueError, match="contiguous"): + c_ext.gen_hashes_numpy(Hasher().hasher, np.repeat(tokens, 2)[::2], TPB, out) + with pytest.raises(ValueError, match="tokens_per_block"): + c_ext.gen_hashes_numpy(Hasher().hasher, tokens, 0, out) + c_ext.gen_hashes_numpy(Hasher().hasher, tokens, TPB, out) # the valid call still works + assert out.tolist() == gen_hashes(tokens, TPB).tolist() diff --git a/tests/test_shm_channel.py b/tests/test_shm_channel.py index 5d4091856..d21577b06 100644 --- a/tests/test_shm_channel.py +++ b/tests/test_shm_channel.py @@ -153,6 +153,60 @@ def test_single_round_trip_local(): ctrl.unlink() +def _explode(): + raise ValueError("poisoned record") + + +class _Poison: + """Pickles fine, fails to unpickle: what a corrupted submit slot looks like.""" + + def __reduce__(self): + return (_explode, ()) + + +def test_undecodable_submit_record_is_dropped_not_fatal(): + """A record the TE cannot decode is dropped with an error and the ring keeps + moving; re-raising on every poll would stall every channel on the node.""" + server_id = f"{SERVER_ID}_poison{os.getpid()}" + _cleanup_shm(server_id, 1) + ctrl = ShmControlBlock(server_id, create=True) + ch = ShmChannel(server_id, 0, create=True) + try: + ch.submit_send(_Poison()) + ch.submit_send("after the poison") + assert ch.submit_recv() == ["after the poison"] + assert ch.submit_recv() == [] # the ring advanced past both + finally: + ch.close() + ch.unlink() + ctrl.close() + ctrl.unlink() + + +def test_every_transfer_type_has_a_wire_index(): + """The TE sends `op.transfer_type.value` for every non-VIRTUAL op, so each + TransferType member must round-trip. LAYERWISE was missing: its KeyError in + the TE's result thread stopped every completion on the node.""" + from flexkv.common.transfer import TransferType + from flexkv.transfer.shm_channel import decode_completed_op, encode_completed_op + for tt in TransferType: + name = None if tt == TransferType.VIRTUAL else tt.value + op = CompletedOp(graph_id=1, op_id=2, transfer_type=name, num_blocks=1, num_bytes=8) + back = decode_completed_op(memoryview(encode_completed_op(op)), 0) + assert back == op, tt + # the member itself is accepted too + op = CompletedOp(graph_id=1, op_id=2, transfer_type=TransferType.LAYERWISE, num_blocks=1) + assert decode_completed_op(memoryview(encode_completed_op(op)), 0).transfer_type == "LAYERWISE" + + +def test_unknown_transfer_type_is_sent_as_none(): + """A name outside the table degrades to None instead of raising.""" + from flexkv.transfer.shm_channel import decode_completed_op, encode_completed_op + op = CompletedOp(graph_id=1, op_id=2, transfer_type="NOT_A_TRANSFER_TYPE", num_blocks=1) + back = decode_completed_op(memoryview(encode_completed_op(op)), 0) + assert back.transfer_type is None and (back.graph_id, back.op_id, back.num_blocks) == (1, 2, 1) + + def test_result_record_carries_durations_and_fits_a_slot(): """The fixed-width record must hold the #297 durations (f64) and still fit the default 64 B result slot; a failed op keeps its flag alongside them.""" From c5ba5ed80c4c7db839cf8a3c799bc97ec272f333 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 14:42:51 +0800 Subject: [PATCH 13/21] radixshmem: drop the shm TE ring; run on FlexKV's own process model (engine mode or KVServer) The radixshmem path had a transport of its own: one node-local TE subprocess (TransferManagerShmTEProcess / te_shm_main) that every DP rank fed over shared-memory submit / result rings (flexkv/transfer/shm_channel.py, shm_channel_handle.py), with per-process op / graph id ranges so the submissions could not collide. That is gone. radixshmem mode now runs on the process model FlexKV already has: - dp_size 1, one engine: the KVTaskEngine lives in the engine process with its TE subprocess (TransferManagerInterProcessHandle), as in the default mode. - dp_size > 1 or FLEXKV_INSTANCE_NUM > 1: server-client mode, one KVServer per node (embedded by the first local DP rank, or started on its own with FLEXKV_SERVER_LAUNCH_MODE=external) whose KVTaskEngine and TE attach the radix-server. server_client_mode is decided exactly as on main again. - shm_radix_bootstrap.adopt_radix_server(model_config, cache_config) attaches with FlexKV's geometry, takes the planned slot counts and the cluster rank over into CacheConfig and detaches. The KVManager runs it before it starts a KVServer or a KVTaskEngine, and KVTaskEngine.__init__ runs it again (idempotent), so an externally started KVServer sizes its pools from the server too. - KVServer.create_server(inherit_env=False) also hands PYTHONPATH, LD_LIBRARY_PATH and PATH to the child: it runs the parent's interpreter and has to import what the parent imports (shmradix, when radixshmem is on PYTHONPATH rather than in site-packages); the embedded server used to die on `import shmradix` there. Removed: flexkv/transfer/shm_channel.py and shm_channel_handle.py, TransferManagerShmTEProcess, shm_te_clears_cuda_visible_devices, the TransferManagerHandle mode="shm", KVManager._init_radix_shmem_path / _spawn_shm_te / _shutdown_radix_shmem_children, the shm_te_* KVTaskEngine parameters, the TransferOp / TransferOpGraph id ranges (flexkv/common/transfer.py is main's again), RadixShmemConfig.te_server_id (default_endpoint replaces the socket derivation), tests/test_shm_channel.py and tests/test_shm_te_cvd.py. docs/radixshmem/config_zh.md and the CHANGELOG describe the process model; the e2e test docstrings follow. The ring design stays reachable at aa26bb5 (branch feat/radixshmem.bak-shm-te-ring). Verified: 455 CPU tests (the CI unit files, the radixshmem suite, the sglang store protocol); the GPU e2e tests/radixshmem/test_e2e_radix_shmem.py passes for dp_size 1 (engine mode) and 2 (KVServer) with byte-exact round trips. --- CHANGELOG.md | 2 +- docs/radixshmem/config_zh.md | 19 +- flexkv/common/radixshmem_config.py | 27 +- flexkv/common/transfer.py | 44 +- flexkv/kvmanager.py | 130 +---- flexkv/kvtask.py | 38 +- flexkv/server/server.py | 8 + flexkv/server/shm_radix_bootstrap.py | 31 +- flexkv/transfer/shm_channel.py | 518 ------------------ flexkv/transfer/shm_channel_handle.py | 378 ------------- flexkv/transfer_manager.py | 117 +--- .../radixshmem/test_e2e_radix_prefetch_p2p.py | 7 +- tests/radixshmem/test_e2e_radix_shmem.py | 24 +- tests/radixshmem/test_radix_shmem_engine.py | 7 +- tests/test_shm_channel.py | 333 ----------- tests/test_shm_te_cvd.py | 17 - 16 files changed, 115 insertions(+), 1585 deletions(-) delete mode 100644 flexkv/transfer/shm_channel.py delete mode 100644 flexkv/transfer/shm_channel_handle.py delete mode 100644 tests/test_shm_channel.py delete mode 100644 tests/test_shm_te_cvd.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 14731053f..ecaadd235 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,7 +14,7 @@ Universal: - radixshmem mode is configured by one small YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `server` (`name`, `endpoint`, `ready_timeout_s`: which radix-server to attach to and how long to wait for it to be reachable and ready) and `client` (`prefetch_timeout_ms`, `prefetch_max_inflight`, `max_outstanding`). No file means `radix-server --name /flexkv` on the local socket. The former `cluster` / `data` / `index` sections are rejected: those settings are radix-server command-line flags now (migration table in `docs/radixshmem/config_zh.md`, launch scripts and a YAML in `examples/radixshmem_configs/`). FlexKV's own TE channel names derive from `server.name`. - radixshmem mode brings no RHT registration chunk of its own any more (the former `REGISTER_CHUNK_TOKENS = 4096` constant, handed to the server as `index.register_chunk_size` in blocks): the chunk is the radix-server's `--register-chunk-tokens` (radixshmem's default 4096 tokens). `RadixGeometry.register_chunk_tokens` forwards a pinned value (0 = the server's), `adopt_geometry` and `CacheEngineRadixShmem.register_chunk_tokens` / `register_chunk_blocks` take the published value over, converted to blocks by radixshmem's rule (`tokens // tokens_per_block`, at least 1); `check_geometry` compares a pinned value and warns when the server's chunk is not a whole number of FlexKV blocks. - radixshmem mode uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. Peer reuse follows the radix-server itself (`world_size > 1`, asked through `shm_radix_bootstrap.radix_server_is_distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). Reference: `docs/radixshmem/config_zh.md`. -- radixshmem mode runs one node-local shm TE process per node (`TransferManagerShmTEProcess`, `te_shm_main`) that every DP rank, and with `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID` every co-located engine, reaches over shared-memory submit / result rings (`flexkv/transfer/shm_channel.py`; the result record carries the worker-measured `wait_ms` / `xfer_ms` / `e2e_ms`). Every `TransferType` value, `LAYERWISE` included, has a wire index and an unknown one degrades to None; a submit record the TE cannot decode is dropped with an error, a graph it cannot submit is reported back as failed, and a delivery error on one channel does not stop the others. Op and graph ids come from per-process ranges (`set_op_id_range` / `set_graph_id_range`, `dp_client_id << 32`), lock-free. +- radixshmem mode keeps FlexKV's process model: with one DP the KVTaskEngine runs in the engine process with its own TE subprocess, otherwise (dp_size > 1, `FLEXKV_INSTANCE_NUM` engines on one node) the DPs are clients of the node's KVServer whose KVTaskEngine and TE attach the radix-server like any other. Wherever a KVTaskEngine is built, `shm_radix_bootstrap.adopt_radix_server` first hands the server FlexKV's geometry and takes its slot counts and cluster rank over into `CacheConfig`. The embedded KVServer child now inherits `PYTHONPATH` / `LD_LIBRARY_PATH` / `PATH` along with the `FLEXKV_*` variables, so it imports what its parent imports (shmradix included) when the packages are not in site-packages. - radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not, and a PUT whose planning fails after its slots were taken returns them and drops the pin before re-raising. `CacheEngineRadixShmem.take` clamps a request to the pool's size instead of letting radixshmem refuse it. - `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`) is trimmed to what `RadixShmemCacheEngine` uses: the `CacheEngineAccel`-compatibility parameters (`device_type`, `evict_ratio`, `evict_start_threshold`, `hit_reward_seconds`, `eviction_policy`, `protected_threshold`, `tokens_per_block=-1`), `take(strict=)`, `match(gpu_matched_blocks=)`, the `mempool` view, `start()`, `store` / `cluster_rank` and the `FLEXKV_TRACE_RADIX_PEER` variable are gone (prefetch logs at debug level; the planner reports mempool metrics itself). - `gen_hashes` and `Hasher.update` hash numpy buffers directly (`c_ext.gen_hashes_numpy` / `update_numpy`) instead of going through `torch.from_numpy`, which is not safe to call concurrently. Hashes are unchanged for int64 tokens; `gen_hashes` keeps requiring int64 and the binding now checks dtype, contiguity and sizes instead of reinterpreting the buffer. diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 202e35376..73cc6b46d 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -22,13 +22,13 @@ radixshmem 侧的接口见 radixshmem 仓库 `python/README.md`。 | 变量 | 默认 | 说明 | |---|---|---| -| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担,KVServer 不启动,每个 DP 进程各建一个 KVTaskEngine 并 attach 同一个 radix-server。在 `flexkv` 首次 import 前设置。 | +| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担。进程模型沿用 FlexKV 原有的:`dp_size=1` 且单实例时 KVTaskEngine 在引擎进程内并自带 TE 子进程;`dp_size>1` 或多实例时走 server-client 模式,每节点一个 KVServer,其中的 KVTaskEngine 和 TE attach radix-server。在 `flexkv` 首次 import 前设置。 | | `FLEXKV_RADIXSHMEM_CONFIG_PATH` | 空 | 第 2 节 YAML 的路径。为空时全部取默认值:attach 本机 `radix-server --name /flexkv`。 | 另有两个 FlexKV 通用变量在该模式下有约束: - `FLEXKV_CPU_LAYOUT` 必须是 `BLOCKFIRST`。一个 SlotStore slot 就是一个连续的 CPU block,LAYERFIRST 给不出这个布局。 -- `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID`:同一节点上多个推理引擎共享同一个 radix-server 和同一个 TE 时用来区分实例(第 4.4 节)。 +- `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID`:同一节点上多个推理引擎共享同一个 radix-server 时用来区分实例,语义与 FlexKV 原有的多实例模式相同(第 4.4 节)。 该模式与 `enable_ssd`、`enable_remote` 互斥,启动时报错。`enable_p2p_cpu` / `enable_p2p_ssd` 也必须为 False: 跨节点复用由 radix-server 自己完成(etcd + RDMA),在它以集群参数启动时自动开启,不经过 FlexKV 的 Redis P2P 路径。 @@ -46,7 +46,7 @@ radixshmem 侧的接口见 radixshmem 仓库 `python/README.md`。 | 键 | 默认 | 说明 | |---|---|---| -| `name` | `/flexkv` | `radix-server --name`,即索引 shm 名。以 `/` 开头。也派生默认 socket 和 FlexKV 自己的 TE channel 前缀(第 5 节)。 | +| `name` | `/flexkv` | `radix-server --name`,即索引 shm 名。以 `/` 开头。也派生默认 socket(第 5 节)。 | | `endpoint` | 空 | gRPC 端点。空为 `unix:///dev/shm/.sock`;server 以 `--endpoint` 改成 TCP 或别的路径时这里写同一个值。 | | `ready_timeout_s` | `600` | 一个 FlexKV 进程等 server **可达且 ready** 的总时长。覆盖运维晚起 server、SlotStore prefault、集群 rendezvous(server 的 `--bootstrap-timeout`)。超时报错并给出启动命令。 | @@ -55,7 +55,7 @@ radixshmem 侧的接口见 radixshmem 仓库 `python/README.md`。 | 键 | 默认 | 说明 | |---|---|---| | `prefetch_timeout_ms` | `5000` | 一次 prefetch 拉取的服务端超时。到期后 job 以本地命中的部分完成。 | -| `prefetch_max_inflight` | `128` | 每个 DP 进程在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | +| `prefetch_max_inflight` | `128` | 每个 KVTaskEngine 在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | | `max_outstanding` | `256` | `RadixClient` 未领取 job 的上限。 | ### 2.3 不再接受的段 @@ -92,12 +92,12 @@ server 收到几何后的规划(radixshmem 的规则):`swa_slots = floor(s `full_slots = (data_bytes − SWA 占用) / full_stride`。任一池算出 0 个 slot、模型有 SWA 而 `--swa-ratio` 为 0, 都在 configure 时拒绝,FlexKV 报 `cannot serve FlexKV's geometry`。 -**采纳**:attach 成功后 `adopt_geometry` 把 `pools.full.num_slots` 写进 `CacheConfig.num_cpu_blocks`, +**采纳**:attach 成功后 `adopt_geometry`(KVManager 和 KVTaskEngine 通过 `adopt_radix_server` 调用)把 `pools.full.num_slots` 写进 `CacheConfig.num_cpu_blocks`, `pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,并带回 server 的 `register_chunk_tokens` 及其 block 数,日志形如 `adopted radix-server /flexkv's geometry: FULL 8605 slots (cpu_cache_gb had given 1524), SWA 1024 slots (had 1024); RHT registration chunk 4096 tokens = 64 blocks`。 之后 TE 的 StorageEngine、cache engine、指标都用采纳后的值。 -**校验**:每个 attach 方(KVManager、cache engine、TE)用 `check_geometry` 复核 server 发布的 +**校验**:每个 attach 方(KVManager、KVTaskEngine、cache engine、TE)用 `check_geometry` 复核 server 发布的 `block_size`、各池 `slot_bytes`、SlotStore stride、SWA 窗口与自己的布局一致,不一致报错退出,不会静默错位传输。 slot 数不在校验范围内,它们是 server 的;`register_chunk_tokens` 只在 FlexKV pin 了值时比对,server 的值不是 `tokens_per_block` 的整数倍时只告警(chunk 取整到整 block)。 @@ -150,14 +150,16 @@ FlexKV 侧每个节点同一份 YAML 即可。`ready_timeout_s` 要不小于 `-- ### 4.4 一节点多引擎共享一个 server -两个独立的推理引擎(各自的 FlexKV、各自的 GPU)attach 同一个 radix-server,互相命中对方存的 KV: +两个独立的推理引擎(各自的 FlexKV、各自的 GPU)attach 同一个 radix-server,互相命中对方存的 KV。这就是 FlexKV +原有的多实例模式: ```bash # 引擎 A # 引擎 B FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=1 ``` -同一份 YAML。`instance 0` 的 dp0 拉起本节点唯一的 TE(`channels = instance_num × dp_size`),TE 等到 +同一份 YAML。`instance_num > 1` 自动进入 server-client 模式:`instance 0` 的 dp0 内嵌本节点唯一的 KVServer(或者用 +`FLEXKV_SERVER_LAUNCH_MODE=external` 单独启动它),KVServer 里的 KVTaskEngine attach radix-server,它的 TE 等到 `instance_num × gpus_per_node` 张 GPU 都注册才 ready,所以两个引擎都要启动。两边模型 / page size / SWA 配置必须相同 (第 3 节)。node-local DP(多机 DP attention)路径下 `local_dp_client_id` 不带 instance,多实例暂不支持。 @@ -171,7 +173,6 @@ FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTA | SlotStore shm | `_data`(`--data-name` 可改) | | gRPC socket | `/dev/shm/.sock`(`--endpoint` 可改;YAML `server.endpoint` 跟着改) | | etcd 键空间 | `radix//...` | -| FlexKV TE channel / ctrl | `/dev/shm/flexkv_te_ch__`、`flexkv_te_ctrl_`,`te_server_id` = `name` 去掉开头的 `/`(`/` 换成 `_`) | --- diff --git a/flexkv/common/radixshmem_config.py b/flexkv/common/radixshmem_config.py index 9991f1586..10a3efe5e 100644 --- a/flexkv/common/radixshmem_config.py +++ b/flexkv/common/radixshmem_config.py @@ -11,9 +11,9 @@ server Which radix-server to attach to: its ``--name`` (which also derives the - default gRPC socket ``unix:///dev/shm/.sock`` and the prefix of - FlexKV's own TE channels), an ``endpoint`` override, and how long a - FlexKV process waits for the server to exist and become ready. + default gRPC socket ``unix:///dev/shm/.sock``), an ``endpoint`` + override, and how long a FlexKV process waits for the server to exist + and become ready. client FlexKV's RadixClient / prefetch settings. @@ -51,12 +51,17 @@ class RadixShmemConfigError(ValueError): """The file is not a valid radixshmem-mode configuration.""" +def default_endpoint(server_name: str) -> str: + """radixshmem's default gRPC socket for ``radix-server --name ``: + ``unix:///dev/shm/ '_'>.sock``.""" + return f"unix:///dev/shm/{server_name.lstrip('/').replace('/', '_')}.sock" + + @dataclasses.dataclass(frozen=True) class RadixServerSettings: """Which radix-server this node's FlexKV attaches to.""" # ``radix-server --name``: the index shm name. Also the default socket - # (``unix:///dev/shm/.sock``) and, sanitized, the prefix of FlexKV's - # TE shm channels on this host (``RadixShmemConfig.te_server_id``). + # (``unix:///dev/shm/.sock``, ``default_endpoint``). name: str = DEFAULT_SERVER_NAME # gRPC endpoint; "" = the default socket derived from ``name``. endpoint: str = "" @@ -100,12 +105,10 @@ def ready_timeout_s(self) -> float: return float(self.server.ready_timeout_s) @property - def te_server_id(self) -> str: - """Prefix of FlexKV's own IPC objects on this host (the TE control - block and channels): the server name without its leading slash, so - two FlexKV deployments on one host that attach to different servers - never share a channel.""" - return self.server.name.lstrip("/").replace("/", "_") + def default_endpoint(self) -> str: + """The gRPC endpoint radixshmem derives from the server name when no + ``endpoint`` is given: ``unix:///dev/shm/.sock``.""" + return default_endpoint(self.server.name) # ------------------------------------------------------------- tests def replace_server(self, **changes: Any) -> "RadixShmemConfig": @@ -117,7 +120,7 @@ def replace_client(self, **changes: Any) -> "RadixShmemConfig": def describe(self) -> str: where = self.path or "(defaults)" return (f"{where}: radix-server {self.server_name} " - f"(endpoint={self.endpoint or 'unix:///dev/shm/' + self.te_server_id + '.sock'}, " + f"(endpoint={self.endpoint or self.default_endpoint}, " f"ready_timeout_s={self.ready_timeout_s:.0f})") diff --git a/flexkv/common/transfer.py b/flexkv/common/transfer.py index 89b824d3b..9d43cdfad 100644 --- a/flexkv/common/transfer.py +++ b/flexkv/common/transfer.py @@ -172,21 +172,6 @@ class TransferOp: # the cost. Kept as a ClassVar so ids stay global across all graphs, which # merge_to_batch_graph relies on when it mixes ops from many tasks. _op_id_counter: ClassVar["itertools.count"] = itertools.count() - # Per-process disjoint range, set by `set_op_id_range()`. Default is the - # full int64 positive range, preserving single-CE behavior. The radix-shmem - # multi-DP path partitions this so 8 CE procs sharing one TE never collide. - # An id is ``start + counter % size``: still lock-free, and it wraps inside - # the range instead of running into a neighbour's. - _op_id_range_start: ClassVar[int] = 0 - _op_id_range_size: ClassVar[int] = 1 << 62 - - @classmethod - def set_op_id_range(cls, start: int, end: int) -> None: - """Restrict generated op_ids to [start, end). Call before any op is - created in this process; it restarts the counter.""" - cls._op_id_range_start = start - cls._op_id_range_size = end - start - cls._op_id_counter = itertools.count() op_id: int = field(init=False) graph_id: int @@ -274,8 +259,7 @@ def __post_init__(self, is_swa: Optional[bool] = None) -> None: raise ValueError(f"src_block_ids and dst_block_ids must have the same number of physical blocks, but got " f"src_block_ids.size={src.size}, " f"dst_block_ids.size={dst.size}") - self.op_id = (TransferOp._op_id_range_start - + next(TransferOp._op_id_counter) % TransferOp._op_id_range_size) + self.op_id = next(TransferOp._op_id_counter) assert src.dtype == _INT64 assert dst.dtype == _INT64 self.valid_block_num = src.size @@ -345,18 +329,9 @@ class TransferOpGraph: # Lock-free for the same reason as TransferOp._op_id_counter: a C-level # __next__ that cannot be interrupted mid-increment. _graph_id_counter = itertools.count() - # Per-process DP-aware range, set via set_graph_id_range(start, end). Used - # so multiple CE processes that share a single TE don't collide on - # graph_id. Default range is the original (0, 2**62), preserving behavior - # in the single-CE path. Same ``start + counter % size`` scheme as - # TransferOp. - _graph_id_range_start = 0 - _graph_id_range_size = 1 << 62 def __init__(self) -> None: - self.graph_id = (TransferOpGraph._graph_id_range_start - + next(TransferOpGraph._graph_id_counter) - % TransferOpGraph._graph_id_range_size) + self.graph_id = next(TransferOpGraph._graph_id_counter) self._op_map: Dict[int, TransferOp] = {} self._ready_ops: Set[int] = set() self._trigger_ops: Set[int] = set() @@ -371,18 +346,7 @@ def __init__(self) -> None: @classmethod def _get_graph_id(cls) -> int: """Kept as the named entry point; __init__ inlines the same counter.""" - return (cls._graph_id_range_start - + next(cls._graph_id_counter) % cls._graph_id_range_size) - - @classmethod - def set_graph_id_range(cls, start: int, end: int) -> None: - """Restrict generated graph_ids to [start, end). Used by the - radix-shmem multi-DP path to give each CE process a disjoint range. - Call before any graph is created in this process; it restarts the - counter.""" - cls._graph_id_range_start = start - cls._graph_id_range_size = end - start - cls._graph_id_counter = itertools.count() + return next(cls._graph_id_counter) def set_graph_id(self, graph_id: int) -> None: self.graph_id = graph_id @@ -1277,7 +1241,7 @@ def merge_to_batch_graph(batch_id: int, put_sinks.append(merged_swa_d2h_op.op_id) if not put_sinks: # No D2H sink: wait for every independent full-KV / SWA leaf - # (H2DISK and/or H2REMOTE). + # (H2DISK and/or H2REMOTE). for op in (merged_h2disk_op, merged_swa_h2disk_op, merged_h2remote_op, merged_swa_h2remote_op): if op is not None: diff --git a/flexkv/kvmanager.py b/flexkv/kvmanager.py index 838043725..e103c6349 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -26,7 +26,6 @@ from flexkv.server.server import KVServer, DPClient from flexkv.kvtask import KVTaskEngine, KVResponse from flexkv.common.config import ModelConfig, CacheConfig, GLOBAL_CONFIG_FROM_ENV, MooncakeTransferEngineConfig -from flexkv.common.transfer import TransferOpGraph from flexkv.integration.dynamo.collector import KVEventCollector from flexkv.common.debug import eviction_log_aggregator, flexkv_logger from flexkv.cache.redis_meta import RedisMeta @@ -86,12 +85,6 @@ def __init__(self, "and enable_p2p_ssd=False (cross-node reuse follows the " "radix-server's cluster flags)" ) - # Prefix of this host's TE channels: the attached radix-server's name. - self._shm_radix_id = None - if self.enable_radixshmem: - from flexkv.common.radixshmem_config import get_radixshmem_config - self._shm_radix_id = get_radixshmem_config().te_server_id - flexkv_logger.info( f"[KVManager] IPC ports: server_recv_port={self.server_recv_port}, " f"gpu_register_port={self.gpu_register_port}" @@ -99,17 +92,13 @@ def __init__(self, ) if self.enable_radixshmem: - flexkv_logger.info(f"[KVManager] radix_shmem is enabled" - f"[KVManager] shm_radix_id: {self._shm_radix_id}") + flexkv_logger.info("[KVManager] radixshmem mode: the CPU tier is the " + "operator's radix-server") # Multi-instance mode also requires server_client_mode - if self.enable_radixshmem: - # Force server_client_mode False — KVServer is bypassed entirely. - self.server_client_mode = False - else: - self.server_client_mode = (model_config.dp_size > 1 or - model_config.instance_num > 1 or - GLOBAL_CONFIG_FROM_ENV.server_client_mode) + self.server_client_mode = (model_config.dp_size > 1 or + model_config.instance_num > 1 or + GLOBAL_CONFIG_FROM_ENV.server_client_mode) self.server_launch_mode = GLOBAL_CONFIG_FROM_ENV.server_launch_mode if self.server_launch_mode not in ("embedded", "external"): raise ValueError( @@ -138,16 +127,22 @@ def __init__(self, self.redis_meta_client = None self.enable_mps = GLOBAL_CONFIG_FROM_ENV.enable_mps self.owns_mps = self.enable_mps and self.server_launch_mode != "external" - # TE-process handle — only the bootstrap process holds this. The - # radix-server itself is the operator's process, not FlexKV's. - self._shm_te_process = None - # Local KVTaskEngine for the radix-shmem path (per-DP). self.kv_task_engine = None self.server_handle = None if self.enable_radixshmem: - self._init_radix_shmem_path(event_collector) - elif self.server_client_mode: + # The CPU tier is the operator's radix-server (``radix-server --name + # --data-bytes ...``); FlexKV never creates one. Hand it + # FlexKV's geometry and take over the slot counts it planned BEFORE + # anything sizes a pool from cache_config: the KVServer or + # KVTaskEngine built below and the TE all read num_cpu_blocks / + # swa.num_slots. The process model is FlexKV's own: engine mode for + # one DP, the KVServer (one per node, shared by FLEXKV_INSTANCE_NUM + # engines) otherwise. + from flexkv.server.shm_radix_bootstrap import adopt_radix_server + adopt_radix_server(model_config, self.cache_config, label="KVManager") + + if self.server_client_mode: # One KVServer per node: with node-local DP the first rank of each # node owns it, so nodes 1..n-1 get their own server instead of # waiting on node 0's. Without node-local DP local_dp_client_id @@ -191,93 +186,6 @@ def __init__(self, event_collector=event_collector, ) - def _init_radix_shmem_path(self, - event_collector: Optional[KVEventCollector]) -> None: - """Initialize the radix-shmem multi-DP path. - - The node's radix-server is a process the operator started - (``radix-server --name --data-bytes ...``); FlexKV never - creates one. Every DP process attaches to it with FlexKV's geometry - (the first one configures the server, the others find it configured), - takes over the slot counts the server planned from its byte budget, - and builds its own KVTaskEngine on it. The single TE subprocess all DP - processes of this node share is spawned by the node-local bootstrap - proc (local DP client 0) only; every other proc attaches to its channel - (``ShmControlBlock.wait_ready``). - - Each CE process gets a disjoint graph/op id range so submissions to the - single shared TE never collide. - """ - from flexkv.common.transfer import TransferOp - from flexkv.server.shm_radix_bootstrap import (adopt_geometry, attach_radix_client, - expected_geometry, radix_cluster_rank) - - TransferOpGraph.set_graph_id_range(self.dp_client_id << 32, - (self.dp_client_id + 1) << 32) - TransferOp.set_op_id_range(self.dp_client_id << 32, - (self.dp_client_id + 1) << 32) - - # Hand the server FlexKV's geometry and take its slot counts BEFORE - # anything sizes a pool from cache_config: the TE's StorageEngine and - # the cache engine both read num_cpu_blocks / swa.num_slots. - geometry = expected_geometry(self.model_config, self.cache_config) - client = attach_radix_client(geometry=geometry, label="KVManager") - try: - adopt_geometry(self.cache_config, client, label="KVManager") - self.cache_config.distributed_node_id = radix_cluster_rank(client) - flexkv_logger.info( - f"[kv manager] radix-server {client.name}: cluster rank " - f"{self.cache_config.distributed_node_id}/{client.info.world_size}, " - f"FlexKV geometry {geometry.describe()}") - finally: - client.close() - - try: - if self.local_dp_client_id == 0: - self._spawn_shm_te() - - # KVTaskEngine reads GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and builds a - # RadixShmemCacheEngine (CPU tier = RadixClient on the radix-server). - self.kv_task_engine = KVTaskEngine( - self.model_config, self.cache_config, - self.gpu_register_port, - redis_meta=self.redis_meta_client, - event_collector=event_collector, - shm_te_server_id=self._shm_radix_id, - shm_te_channel_id=self.local_dp_client_id, - ) - except BaseException: - # A failure after the TE subprocess was spawned must not leave it - # running (it would wait for GPU registrations forever). - self._shutdown_radix_shmem_children() - raise - - def _shutdown_radix_shmem_children(self) -> None: - # getattr: shutdown() must work on a KVManager whose __init__ did not - # get this far (or was bypassed, as the unit tests do). - te = getattr(self, "_shm_te_process", None) - if te is not None: - self._shm_te_process = None - te.shutdown() - - def _spawn_shm_te(self) -> None: - """Bootstrap proc (local dp 0) only: spawn the TE subprocess every DP - process of this node feeds over its shm channel. The TE attaches to - the same radix-server (its SlotStore is the CPU pool) with the cache - config whose slot counts were just adopted.""" - from flexkv.transfer_manager import TransferManagerShmTEProcess - - total_clients = self.model_config.total_clients - if self.model_config.local_dp_size is not None: - total_clients = self.model_config.local_dp_size - self._shm_te_process = TransferManagerShmTEProcess( - self.model_config, self.cache_config, - gpu_register_port=self.gpu_register_port, - server_id=self._shm_radix_id, - num_channels=total_clients, - ) - self._shm_te_process.start() - def start(self) -> None: if self.owns_mps: # try to start MPS @@ -312,10 +220,6 @@ def shutdown(self) -> None: if self.kv_task_engine is not None: self.kv_task_engine.shutdown() - # Multi-DP radix-shmem teardown — only the bootstrap proc owns these. - # TE first: its workers map the server's SlotStore. - self._shutdown_radix_shmem_children() - if self.owns_mps: flexkv_logger.info( "MPS is enabled. To stop MPS daemon manually, run: " diff --git a/flexkv/kvtask.py b/flexkv/kvtask.py index bcaa3cfb9..73e78cab8 100644 --- a/flexkv/kvtask.py +++ b/flexkv/kvtask.py @@ -145,8 +145,6 @@ def __init__(self, gpu_register_port: Optional[str] = None, redis_meta: RedisMeta = None, event_collector: Optional[KVEventCollector] = None, - shm_te_server_id: Optional[str] = None, - shm_te_channel_id: Optional[int] = None, ): if not cache_config.enable_cpu: raise ValueError("enable_cpu must be True") @@ -179,30 +177,20 @@ def __init__(self, self.prefetch_jobs: Dict[int, Any] = {} if GLOBAL_CONFIG_FROM_ENV.enable_radixshmem: # The CPU tier is a radix-server (shared index + SlotStore); its - # planners are a GlobalCacheEngine subclass. + # planners are a GlobalCacheEngine subclass. Take the server's slot + # counts over first: the planner below and the TE size their pools + # from cache_config. Idempotent; the KVManager did it too, but a + # KVServer started on its own (FLEXKV_SERVER_LAUNCH_MODE=external) + # arrives here with its own config. from flexkv.cache.radix_shmem_planner import RadixShmemCacheEngine + from flexkv.server.shm_radix_bootstrap import adopt_radix_server + adopt_radix_server(model_config, cache_config, label="KVTaskEngine") self.cache_engine = RadixShmemCacheEngine( cache_config, model_config, redis_meta, event_collector) else: self.cache_engine = GlobalCacheEngine(cache_config, model_config, redis_meta, event_collector) - # Multi-DP shm path: connect this CE to a pre-existing TE process - # via a named ShmChannel rather than spawning a new TE subprocess. - use_shm_te = (shm_te_server_id is not None - and shm_te_channel_id is not None) - if use_shm_te and not self.model_config.use_trtllm_subprocess: - self.transfer_handles = [TransferManagerHandle( - # Left behind by a rename: the sibling "process" branch below - # passes `model_config`, and no *_for_transfer variant exists — - # so the shm-TE path (radix_shmem) NameError'd on first use. - model_config, - self.cache_config, - mode="shm", - gpu_register_port=gpu_register_port, - shm_server_id=shm_te_server_id, - shm_channel_id=shm_te_channel_id, - )] - elif not self.model_config.use_trtllm_subprocess: + if not self.model_config.use_trtllm_subprocess: self.transfer_handles = [TransferManagerHandle( model_config, cache_config, @@ -231,7 +219,8 @@ def __init__(self, ] self.transfer_handles[0]._handle.send_config_to_remotes() - # A node-local shm TE replaces the legacy cross-node remote manager. + # Node-local DP: every node runs its own KVServer and TE, so the + # cross-node remote transfer manager is not needed. needs_remote_transfer_manager = self.model_config.local_dp_size is None if self.model_config.nnodes > 1 and needs_remote_transfer_manager: # Bind the handle rather than reading it back with a negative @@ -1116,13 +1105,8 @@ def __init__(self, gpu_register_port: Optional[str] = None, redis_meta: Optional[RedisMeta] = None, event_collector: Optional[KVEventCollector] = None, - shm_te_server_id: Optional[str] = None, - shm_te_channel_id: Optional[int] = None, ): - super().__init__(model_config, cache_config, gpu_register_port, - redis_meta, event_collector, - shm_te_server_id=shm_te_server_id, - shm_te_channel_id=shm_te_channel_id) + super().__init__(model_config, cache_config, gpu_register_port, redis_meta, event_collector) self.tracer = FlexKVTracer() self.tracer.trace_config(model_config, cache_config, gpu_layout=None) diff --git a/flexkv/server/server.py b/flexkv/server/server.py index 438f14740..19176f29f 100644 --- a/flexkv/server/server.py +++ b/flexkv/server/server.py @@ -256,6 +256,14 @@ def create_server(cls, for key, val in os.environ.items(): if key.startswith("FLEXKV_") and key not in env: env[key] = val + # The child runs the parent's interpreter and must import what + # the parent imports (flexkv, shmradix in radixshmem mode) and + # find the same shared libraries; these are the variables that + # locate them when the packages are not installed into + # site-packages. + for key in ("PYTHONPATH", "LD_LIBRARY_PATH", "PATH"): + if key in os.environ and key not in env: + env[key] = os.environ[key] cvd = os.environ.get('CUDA_VISIBLE_DEVICES') if cvd is not None and 'CUDA_VISIBLE_DEVICES' not in env: diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index d98fba10d..642676668 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -41,7 +41,8 @@ from flexkv.common.config import (GLOBAL_CONFIG_FROM_ENV, CacheConfig, LayerGroupSpec, ModelConfig, SWAPoolConfig) from flexkv.common.debug import flexkv_logger -from flexkv.common.radixshmem_config import RadixShmemConfig, get_radixshmem_config +from flexkv.common.radixshmem_config import (RadixShmemConfig, default_endpoint, + get_radixshmem_config) from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType try: @@ -269,7 +270,7 @@ def attach_radix_client(name: Optional[str] = None, if max_outstanding is None: max_outstanding = rcfg.client.max_outstanding spec = geometry.to_shmradix() if isinstance(geometry, RadixGeometry) else geometry - where = endpoint or f"unix:///dev/shm/{name.lstrip('/').replace('/', '_')}.sock" + where = endpoint or default_endpoint(name) deadline = time.monotonic() + float(timeout_s) last: Optional[BaseException] = None @@ -437,6 +438,32 @@ def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", return counts +def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, + *, + rcfg: Optional[RadixShmemConfig] = None, + label: str = "radixshmem") -> Dict[str, int]: + """Attach to this node's radix-server with FlexKV's geometry, take over the + slot counts it planned (:func:`adopt_geometry`) and its cluster rank + (``cache_config.distributed_node_id``), then detach. Every process that + sizes something from ``cache_config`` runs this before it does: the + KVManager before it starts a KVServer or a KVTaskEngine, the KVTaskEngine + host itself (a KVServer may be started on its own) and, with a live client, + the TE. Idempotent: the server accepts the same geometry any number of + times.""" + geometry = expected_geometry(model_config, cache_config) + client = attach_radix_client(rcfg=rcfg, geometry=geometry, label=label) + try: + counts = adopt_geometry(cache_config, client, label=label) + cache_config.distributed_node_id = radix_cluster_rank(client) + flexkv_logger.info( + f"{label}: radix-server {client.name}: cluster rank " + f"{cache_config.distributed_node_id}/{client.info.world_size}, " + f"FlexKV geometry {geometry.describe()}") + return counts + finally: + client.close() + + def radix_server_is_distributed(rcfg: Optional[RadixShmemConfig] = None, *, timeout_s: Optional[float] = None, diff --git a/flexkv/transfer/shm_channel.py b/flexkv/transfer/shm_channel.py deleted file mode 100644 index efa12170c..000000000 --- a/flexkv/transfer/shm_channel.py +++ /dev/null @@ -1,518 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# cython: boundscheck=True, wraparound=True -""" -Shared-memory IPC channel for CacheEngine ↔ TransferEngine communication in -multi-DP FlexKV. - -Adapted from PR #144 (commit 5e262ca, originally for DPClient ↔ KVServer). -Slimmed down for the CE↔TE use case: - - Per-channel size is small (256 KB ring + 256 KB sync) because the only - payloads are pickled TransferOpGraph (submit) and CompletedOp lists (wait). - - Sync request/response slot is dropped — CE↔TE is fully fire-and-forget on - both directions; submit is async, completions are pushed asynchronously. - - Two SPSC ring buffers per channel: `submit` (CE→TE) and `result` (TE→CE). - Each side futex-waits on its own counter when the ring is empty. - -Layout per channel (one /dev/shm file per CE): - [0..64) submit_write_pos (uint64, CE writes) - [64..128) submit_read_pos (uint64, TE writes) - [128..192) submit_wake (int32, CE bumps + futex_wake) - [192..256) result_write_pos (uint64, TE writes) - [256..320) result_read_pos (uint64, CE writes) - [320..384) result_wake (int32, TE bumps + futex_wake) - [384..ring_off) reserved - [ring_off ..) submit ring (slot * SUBMIT_SLOTS) - result ring (slot * RESULT_SLOTS) - -A single `ShmControlBlock` (separate /dev/shm file) carries a global wake -counter that the TE polls when it has more than one channel attached, so it -can sleep idly without per-channel futex wait. -""" -from __future__ import annotations - -import ctypes -import ctypes.util -import mmap -import os -import pickle -import platform -import struct -from typing import Any, List, Optional - -# ── Linux futex wrappers ──────────────────────────────────────────────── - -_libc = ctypes.CDLL(ctypes.util.find_library("c"), use_errno=True) - -# futex syscall number is arch-specific; pick at import time. -_MACHINE = platform.machine().lower() -if _MACHINE in ("x86_64", "amd64"): - _SYS_FUTEX = 202 -elif _MACHINE in ("aarch64", "arm64"): - _SYS_FUTEX = 98 -else: # pragma: no cover - raise RuntimeError(f"Unsupported machine for futex syscall: {_MACHINE}") -_FUTEX_WAIT = 0 -_FUTEX_WAKE = 1 - - -def _futex_wait(addr: int, expected: int, timeout_ns: Optional[int] = None) -> int: - if timeout_ns is None: - return _libc.syscall( - _SYS_FUTEX, ctypes.c_void_p(addr), - _FUTEX_WAIT, ctypes.c_int(expected), - ctypes.c_void_p(0), ctypes.c_void_p(0), ctypes.c_int(0), - ) - # struct timespec - ts = (ctypes.c_long * 2)(timeout_ns // 1_000_000_000, - timeout_ns % 1_000_000_000) - return _libc.syscall( - _SYS_FUTEX, ctypes.c_void_p(addr), - _FUTEX_WAIT, ctypes.c_int(expected), - ctypes.byref(ts), ctypes.c_void_p(0), ctypes.c_int(0), - ) - - -def _futex_wake(addr: int, count: int = 1) -> int: - return _libc.syscall( - _SYS_FUTEX, ctypes.c_void_p(addr), - _FUTEX_WAKE, ctypes.c_int(count), - ctypes.c_void_p(0), ctypes.c_void_p(0), ctypes.c_int(0), - ) - - -# ── Layout constants ──────────────────────────────────────────────────── - -_CL = 64 # cache line - -# Header lives in the first 6 cache lines; ring data starts on a page boundary. -OFF_SUBMIT_W = 0 * _CL -OFF_SUBMIT_R = 1 * _CL -OFF_SUBMIT_WAKE = 2 * _CL -OFF_RESULT_W = 3 * _CL -OFF_RESULT_R = 4 * _CL -OFF_RESULT_WAKE = 5 * _CL -HEADER_SIZE = 6 * _CL # 384 B - -# Submit ring holds pickled TransferOpGraphs, fragmented across slots so a payload -# larger than one slot spans several (32 KB × 8192 = 256 MB, ~256 MB max message). -# Fragmentation retired the old "must fit one slot" constraint that made high-QPS -# as_batch=True graphs overflow (observed 152 KB @2500 with prefetch). -DEFAULT_SUBMIT_SLOTS = 8192 # power of 2 -DEFAULT_SUBMIT_SLOT_SIZE = 32 * 1024 -DEFAULT_SLOT_SIZE = DEFAULT_SUBMIT_SLOT_SIZE # back-compat alias for `slot_size=` - -# Result ring holds one fixed-width CompletedOp record per slot (64 B × 65536 = 4 MB). -DEFAULT_RESULT_SLOTS = 65536 # power of 2 -DEFAULT_RESULT_SLOT_SIZE = 64 # cache line; one 54 B CompletedOp record - -_PAGE = 4096 - - -def _round_up(x: int, m: int) -> int: - return (x + m - 1) // m * m - - -# Submit-ring fragment header (per slot): payload bytes in this slot (u32) + a -# last-fragment flag (u8). A message is the concatenation of fragments up to and -# including the one with is_last=1. -_FRAG_HDR = struct.Struct(" int: - """Wire index of a CompletedOp.transfer_type (a TransferType value, the - member itself, or None). A name the table does not know goes out as None - with one warning per name: the CE loses that op's type attribution, which - beats a KeyError in the TE's result thread that would stop every - completion on the node.""" - if tt is None: - return _TT_NONE - key = getattr(tt, "value", tt) - idx = _TT_NAME_TO_IDX.get(key) - if idx is None: - if key not in _tt_unknown_warned: - _tt_unknown_warned.add(key) - from flexkv.common.debug import flexkv_logger - flexkv_logger.warning( - f"shm channel: transfer type {key!r} has no wire index; it is sent as " - f"None (add it to shm_channel._TT_NAMES)") - return _TT_NONE - return idx - - -def encode_completed_op(op: Any) -> bytes: - """Pack a CompletedOp into its fixed-width record.""" - tt_idx = transfer_type_index(op.transfer_type) - flags = 1 if getattr(op, "failed", False) else 0 - return _COMPLETED_OP.pack( - op.graph_id, op.op_id, tt_idx, op.num_blocks, op.num_bytes, flags, - float(getattr(op, "wait_ms", 0.0)), - float(getattr(op, "xfer_ms", 0.0)), - float(getattr(op, "e2e_ms", 0.0)), - ) - - -def decode_completed_op(buf: Any, off: int) -> Any: - """Unpack a CompletedOp record from `buf` at byte offset `off`.""" - from flexkv.common.transfer import CompletedOp - graph_id, op_id, tt_idx, num_blocks, num_bytes, flags, wait_ms, xfer_ms, e2e_ms = \ - _COMPLETED_OP.unpack_from(buf, off) - tt = None if tt_idx == _TT_NONE or tt_idx >= len(_TT_NAMES) else _TT_NAMES[tt_idx] - return CompletedOp( - graph_id=graph_id, - op_id=op_id, - transfer_type=tt, - num_blocks=num_blocks, - num_bytes=num_bytes, - wait_ms=wait_ms, - xfer_ms=xfer_ms, - e2e_ms=e2e_ms, - failed=bool(flags & 1), - ) - - -# ── ShmControlBlock ───────────────────────────────────────────────────── - -CTRL_WAKE = 0 -CTRL_READY = _CL -CTRL_SIZE = _PAGE - - -def _safe_id(server_id: str) -> str: - return server_id.replace("/", "_").replace(":", "_").strip("_") - - -class ShmControlBlock: - """Optional global wake counter used by the TE when polling N channels.""" - - def __init__(self, server_id: str, create: bool = False): - self.server_id = server_id - safe = _safe_id(server_id) - self.shm_path = f"/dev/shm/flexkv_te_ctrl_{safe}" - - if create: - fd = os.open(self.shm_path, os.O_CREAT | os.O_RDWR, 0o666) - os.ftruncate(fd, CTRL_SIZE) - self.buf = mmap.mmap(fd, CTRL_SIZE) - os.close(fd) - self.buf[:] = b"\x00" * CTRL_SIZE - else: - fd = os.open(self.shm_path, os.O_RDWR) - self.buf = mmap.mmap(fd, CTRL_SIZE) - os.close(fd) - - self._base = ctypes.addressof(ctypes.c_char.from_buffer(self.buf)) - self._wake = ctypes.c_int32.from_address(self._base + CTRL_WAKE) - self._ready = ctypes.c_int32.from_address(self._base + CTRL_READY) - - @property - def _wake_addr(self) -> int: - return self._base + CTRL_WAKE - - @property - def _ready_addr(self) -> int: - return self._base + CTRL_READY - - def notify(self) -> None: - # Read-modify-write is not atomic across processes; the TE uses snapshot - # comparison so any change wakes it. - self._wake.value += 1 - _futex_wake(self._wake_addr, 1) - - def get_wake(self) -> int: - return self._wake.value - - def wait(self, expected: int, timeout_ns: Optional[int] = None) -> None: - _futex_wait(self._wake_addr, expected, timeout_ns) - - def set_ready(self) -> None: - self._ready.value = 1 - _futex_wake(self._ready_addr, 0x7FFFFFFF) - - def wait_ready(self, timeout_s: float = 60.0) -> bool: - import time - deadline = time.monotonic() + timeout_s - while time.monotonic() < deadline: - if self._ready.value != 0: - return True - _futex_wait(self._ready_addr, 0, - timeout_ns=int(0.5 * 1_000_000_000)) - return False - - def close(self) -> None: - if self.buf is not None: - self.buf.close() - self.buf = None - - def unlink(self) -> None: - try: - os.unlink(self.shm_path) - except FileNotFoundError: - pass - - -# ── ShmChannel (CE ↔ TE) ──────────────────────────────────────────────── - -class ShmChannel: - """One bi-directional channel between a single CE and the TE. - - Two SPSC rings: CE→TE submit, TE→CE result. Each side futex-waits on its - own wake counter when its consumer ring is empty. - """ - - def __init__(self, - server_id: str, - channel_id: int, - create: bool = False, - submit_slots: int = DEFAULT_SUBMIT_SLOTS, - result_slots: int = DEFAULT_RESULT_SLOTS, - slot_size: int = DEFAULT_SUBMIT_SLOT_SIZE, - result_slot_size: int = DEFAULT_RESULT_SLOT_SIZE): - assert submit_slots & (submit_slots - 1) == 0, \ - "submit_slots must be power of 2" - assert result_slots & (result_slots - 1) == 0, \ - "result_slots must be power of 2" - assert result_slot_size >= COMPLETED_OP_WIRE_SIZE, \ - f"result_slot_size {result_slot_size} < CompletedOp record " \ - f"{COMPLETED_OP_WIRE_SIZE}" - - self.channel_id = channel_id - self.submit_slots = submit_slots - self.result_slots = result_slots - self.slot_size = slot_size # submit ring; result ring uses result_slot_size - self.result_slot_size = result_slot_size - - safe = _safe_id(server_id) - self.shm_path = f"/dev/shm/flexkv_te_ch_{safe}_{channel_id}" - - # Lay out: header -> aligned to page -> submit ring -> result ring. - self._submit_off = _round_up(HEADER_SIZE, _PAGE) - self._result_off = self._submit_off + submit_slots * slot_size - total = self._result_off + result_slots * result_slot_size - - self.total_size = total - - if create: - fd = os.open(self.shm_path, os.O_CREAT | os.O_RDWR, 0o666) - os.ftruncate(fd, total) - self.buf = mmap.mmap(fd, total) - os.close(fd) - self.buf[:HEADER_SIZE] = b"\x00" * HEADER_SIZE - else: - fd = os.open(self.shm_path, os.O_RDWR) - self.buf = mmap.mmap(fd, total) - os.close(fd) - - self._base = ctypes.addressof(ctypes.c_char.from_buffer(self.buf)) - self._submit_w = ctypes.c_uint64.from_address(self._base + OFF_SUBMIT_W) - self._submit_r = ctypes.c_uint64.from_address(self._base + OFF_SUBMIT_R) - self._submit_wake = ctypes.c_int32.from_address(self._base + OFF_SUBMIT_WAKE) - self._result_w = ctypes.c_uint64.from_address(self._base + OFF_RESULT_W) - self._result_r = ctypes.c_uint64.from_address(self._base + OFF_RESULT_R) - self._result_wake = ctypes.c_int32.from_address(self._base + OFF_RESULT_WAKE) - - # ---- futex helpers ---- - - @property - def _submit_wake_addr(self) -> int: - return self._base + OFF_SUBMIT_WAKE - - @property - def _result_wake_addr(self) -> int: - return self._base + OFF_RESULT_WAKE - - def _bump_wake(self, ptr: ctypes.c_int32, addr: int) -> None: - ptr.value += 1 - _futex_wake(addr, 1) - - # ---- ring helpers ---- - - def _ring_full(self, w: int, r: int, slots: int) -> bool: - return ((w + 1) & (slots - 1)) == r - - def _ring_used(self, w: int, r: int, slots: int) -> int: - return (w - r) & (slots - 1) - - # ---- CE side: submit + recv result ---- - - def submit_send(self, payload: Any) -> None: - """Enqueue a payload to TE, fragmenting it across slots if it exceeds one. - - All fragments are written first, then submit_w is advanced once, so the TE - never observes a partial message. Spins+yields until enough contiguous - slots are free.""" - blob = pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL) - slots = self.submit_slots - body = self.slot_size - _FRAG_HDR_SIZE - nfrag = max(1, (len(blob) + body - 1) // body) - if nfrag > slots - 1: - raise ValueError( - f"payload needs {nfrag} fragments but submit ring holds " - f"{slots - 1}; raise submit_slots or slot_size") - - wp = self._submit_w.value - spin = 0 - while self._ring_used(wp, self._submit_r.value, slots) + nfrag > slots - 1: - if spin > 1_000_000: - raise RuntimeError("shm channel submit ring full") - if spin > 1000: - os.sched_yield() - spin += 1 - - w = wp - for i in range(nfrag): - chunk = blob[i * body:(i + 1) * body] - off = self._submit_off + w * self.slot_size - self.buf[off:off + _FRAG_HDR_SIZE] = _FRAG_HDR.pack( - len(chunk), 1 if i == nfrag - 1 else 0) - self.buf[off + _FRAG_HDR_SIZE:off + _FRAG_HDR_SIZE + len(chunk)] = chunk - w = (w + 1) & (slots - 1) - self._submit_w.value = w - self._bump_wake(self._submit_wake, self._submit_wake_addr) - - def result_recv(self, timeout_s: Optional[float] = None) -> List[Any]: - """Drain pending TE→CE results, one CompletedOp per slot. Blocks up to - `timeout_s` if empty.""" - out: List[Any] = [] - rp = self._result_r.value - wp = self._result_w.value - slots = self.result_slots - if rp == wp and timeout_s is not None and timeout_s > 0: - wake = self._result_wake.value - # Re-check; TE might have arrived between read and wait. - wp = self._result_w.value - if rp == wp: - if timeout_s == float("inf"): - _futex_wait(self._result_wake_addr, wake) - else: - _futex_wait(self._result_wake_addr, wake, - timeout_ns=int(timeout_s * 1_000_000_000)) - wp = self._result_w.value - - while rp != wp: - out.append(decode_completed_op( - self.buf, self._result_off + rp * self.result_slot_size)) - rp = (rp + 1) & (slots - 1) - if out: - self._result_r.value = rp - return out - - # ---- TE side: recv submit + send result ---- - - def submit_recv(self) -> List[Any]: - """Drain CE→TE submissions (non-blocking), reassembling fragmented - messages. submit_r is advanced only past fully-received messages.""" - out: List[Any] = [] - rp = self._submit_r.value - wp = self._submit_w.value - slots = self.submit_slots - frags: List[bytes] = [] - while rp != wp: - off = self._submit_off + rp * self.slot_size - n, is_last = _FRAG_HDR.unpack_from(self.buf, off) - frags.append(bytes(self.buf[off + _FRAG_HDR_SIZE: - off + _FRAG_HDR_SIZE + n])) - rp = (rp + 1) & (slots - 1) - if is_last: - payload = b"".join(frags) - frags.clear() - self._submit_r.value = rp # release this message's slots (payload is a copy) - try: - out.append(pickle.loads(payload)) - except Exception as e: # noqa: BLE001 - a record we cannot decode - # Dropping it costs the sender one task; keeping it would - # re-raise on every poll and stall the whole ring. - from flexkv.common.debug import flexkv_logger - flexkv_logger.error( - f"shm channel {self.channel_id}: dropping an undecodable " - f"{len(payload)}-byte submit record ({e!r})") - return out - - def result_send(self, ops: List[Any]) -> None: - """Enqueue a batch of CompletedOps, one fixed-width record per slot, and - wake the CE once at the end. Spins if the ring fills rather than dropping a - completion (which would hang the owning task); with 65536 slots that is - effectively unreachable.""" - if not ops: - return - slots = self.result_slots - slot_sz = self.result_slot_size - base = self._result_off - wp = self._result_w.value - warned = False - for op in ops: - if self._ring_full(wp, self._result_r.value, slots): - # Full mid-batch: publish+wake so the CE drains, then spin. - self._bump_wake(self._result_wake, self._result_wake_addr) - spin = 0 - while self._ring_full(wp, self._result_r.value, slots): - if spin > 1000: - os.sched_yield() - spin += 1 - if spin % 5_000_000 == 0 and not warned: - try: - from flexkv.common.debug import flexkv_logger - flexkv_logger.error( - f"shm channel result ring stuck full " - f"(slots={slots}); is the CE consumer alive?" - ) - except Exception: - pass - warned = True - off = base + wp * slot_sz - self.buf[off:off + COMPLETED_OP_WIRE_SIZE] = encode_completed_op(op) - wp = (wp + 1) & (slots - 1) - self._result_w.value = wp - self._bump_wake(self._result_wake, self._result_wake_addr) - - # ---- Submit-wake fileno: lets TE selector wait on this channel ---- - - @property - def submit_wake_fd(self) -> int: - # We don't have a real eventfd; selector users should poll get_wake() - # delta + futex_wait via ShmControlBlock. Returning -1 signals "no fd". - return -1 - - # ---- Lifecycle ---- - - def close(self) -> None: - if self.buf is not None: - self.buf.close() - self.buf = None - - def unlink(self) -> None: - try: - os.unlink(self.shm_path) - except FileNotFoundError: - pass - - def __del__(self) -> None: - try: - self.close() - except Exception: - pass diff --git a/flexkv/transfer/shm_channel_handle.py b/flexkv/transfer/shm_channel_handle.py deleted file mode 100644 index 57af4c17b..000000000 --- a/flexkv/transfer/shm_channel_handle.py +++ /dev/null @@ -1,378 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# cython: boundscheck=True, wraparound=True -""" -Shared-memory variant of TransferManagerHandle for the multi-DP path. - -Architecture: -- N CE processes each hold one `TransferManagerShmChannelHandle` connected to - the single TE process via a `ShmChannel` named after the (server_id, - channel_id). -- The TE process runs a multi-channel dispatcher loop (`_te_shm_main`) that - polls all N submit rings, hands graphs to the underlying TransferManager, - and routes each completed op back to its originating channel via a - `graph_id → channel_id` map. -""" -from __future__ import annotations - -import os -import queue -import threading -import time -from typing import Dict, List, Optional, Tuple - -import nvtx - -from flexkv.common.config import CacheConfig, ModelConfig -from flexkv.common.debug import flexkv_logger -from flexkv.common.transfer import CompletedOp, TransferOpGraph -from flexkv.transfer.shm_channel import ShmChannel, ShmControlBlock - - -# Wire format: small dicts so we don't pay for full TransferOpGraph re-pickling -# more than necessary. Submit messages carry the graph itself; we'll let pickle -# handle the structure. - -class _SubmitMsg: - __slots__ = ("graph", "task_end_op_id", "is_batch") - - def __init__(self, graph, task_end_op_id: int = -1, is_batch: bool = False): - self.graph = graph - self.task_end_op_id = task_end_op_id - self.is_batch = is_batch - - -# The result ring carries fixed-width CompletedOp records directly (no wrapper). - - -# CE-side handle --------------------------------------------------------- - -class TransferManagerShmChannelHandle: - """CE-side handle that submits transfer graphs to a shared TE via shmem.""" - - def __init__(self, - model_config: ModelConfig, - cache_config: CacheConfig, - server_id: str, - channel_id: int, - file_wait_timeout_s: float = 60.0): - from flexkv.transfer.shm_channel import _safe_id - - self.model_config = model_config - self.cache_config = cache_config - self.server_id = server_id - self.channel_id = channel_id - - safe = _safe_id(server_id) - self._ctrl_path = f"/dev/shm/flexkv_te_ctrl_{safe}" - self._ch_path = f"/dev/shm/flexkv_te_ch_{safe}_{channel_id}" - self._file_wait_timeout_s = file_wait_timeout_s - - # Poll for shm files (created in TE subprocess by setup_channels()). - # Existence does NOT mean the TM is initialized — that's signalled - # by the ctrl ready flag and observed via `is_ready()`. - deadline = time.time() + file_wait_timeout_s - while time.time() < deadline: - if os.path.exists(self._ctrl_path) and os.path.exists(self._ch_path): - break - time.sleep(0.05) - else: - raise RuntimeError( - f"Timed out waiting for TE shm files: " - f"{self._ctrl_path}, {self._ch_path}" - ) - # Attaching is non-blocking once files exist. - self._ctrl = ShmControlBlock(server_id, create=False) - self._channel = ShmChannel(server_id, channel_id, create=False) - - # TransferManagerHandleBase interface ------------------------------- - - def start(self) -> None: - # Nothing to do — TE creates and starts the channel. - pass - - def is_ready(self) -> bool: - # The TE subprocess flips the ctrl ready flag after the - # TransferManager (incl. GPU registration) finishes initializing. - return self._ctrl._ready.value != 0 - - def submit(self, transfer_graph: TransferOpGraph, - task_end_op_id: int = -1) -> None: - nvtx_range = nvtx.start_range( - message="TransferManagerShmChannelHandle.submit", color="green" - ) - self._channel.submit_send(_SubmitMsg(transfer_graph, task_end_op_id)) - self._ctrl.notify() # wake the TE: an idle _poll_submits parks on the ctrl futex for up to 100 ms - nvtx.end_range(nvtx_range) - - def submit_batch(self, transfer_graphs: List[TransferOpGraph]) -> None: - # Send each graph as its own submit message — keeps the TE side - # simple. Could be batched into a list for fewer pickle calls if - # benchmarks show it matters. - for g in transfer_graphs: - self._channel.submit_send(_SubmitMsg(g, -1, is_batch=True)) - if transfer_graphs: - self._ctrl.notify() # see submit() - - def wait(self, timeout: Optional[float] = None) -> List[CompletedOp]: - if timeout is None: - timeout = 0.0 - out: List[CompletedOp] = self._channel.result_recv(timeout_s=timeout) - if out and os.environ.get("FLEXKV_TRACE_TE", "0") == "1": - completed_graphs = sorted({op.graph_id for op in out - if op.is_graph_completed()}) - flexkv_logger.info( - f"[TE-TRACE] CE recv ch={self.channel_id} " - f"completed_graphs={completed_graphs} n_ops={len(out)}" - ) - return out - - def shutdown(self) -> None: - try: - self._channel.close() - except Exception: - pass - try: - self._ctrl.close() - except Exception: - pass - - def __del__(self) -> None: - try: - self.shutdown() - except Exception: - pass - - -# TE-side multi-channel dispatcher -------------------------------------- - -class _TEShmDispatcher: - """Polls N channels, forwards submits to TransferManager, routes results. - - Two-phase startup: - 1) `setup_channels()` creates the control block + per-channel shm files - and sets the ready flag. CE-side handles can attach as soon as this - returns. TM does not need to exist yet. - 2) `start_dispatch(transfer_manager)` launches the polling threads. Call - this after the TM has finished initializing. - """ - - def __init__(self, server_id: str, num_channels: int): - self._tm = None - self._server_id = server_id - self._num_channels = num_channels - self._ctrl: Optional[ShmControlBlock] = None - self._channels: List[ShmChannel] = [] - # graph_id -> channel_id (submitter) - self._graph_owner: Dict[int, int] = {} - self._owner_lock = threading.Lock() - self._stop = threading.Event() - self._poll_thread: Optional[threading.Thread] = None - self._result_thread: Optional[threading.Thread] = None - - def setup_channels(self) -> None: - """Create shm control block + per-channel files. Idempotent w.r.t. CE - attaches — CE-side handles only need these files to exist.""" - self._ctrl = ShmControlBlock(self._server_id, create=True) - self._channels = [ - ShmChannel(self._server_id, ch_id, create=True) - for ch_id in range(self._num_channels) - ] - flexkv_logger.info( - f"TE shm dispatcher: {self._num_channels} channels created on " - f"server_id={self._server_id}" - ) - - def start_dispatch(self, transfer_manager) -> None: - """Bind the TransferManager and start the polling threads. The ctrl - ready flag is flipped so CEs that were spinning on `wait_ready` - can proceed.""" - self._tm = transfer_manager - assert self._ctrl is not None, "setup_channels() must run first" - self._ctrl.set_ready() - self._poll_thread = threading.Thread( - target=self._poll_submits, daemon=True, name="te-shm-poll" - ) - self._result_thread = threading.Thread( - target=self._poll_results, daemon=True, name="te-shm-result" - ) - self._poll_thread.start() - self._result_thread.start() - - def shutdown(self) -> None: - self._stop.set() - # Wake up futex waiters so threads can exit. - if self._ctrl is not None: - self._ctrl.notify() - for ch in self._channels: - try: - ch.close() - except Exception: - pass - self._channels = [] - if self._ctrl is not None: - try: - self._ctrl.close() - self._ctrl.unlink() - except Exception: - pass - self._ctrl = None - - def _poll_submits(self) -> None: - idle_spins = 0 - while not self._stop.is_set(): - had_work = False - for ch in self._channels: - msgs = ch.submit_recv() - if not msgs: - continue - had_work = True - for m in msgs: - if not isinstance(m, _SubmitMsg): - flexkv_logger.warning( - f"TE got unexpected submit msg type: {type(m)}" - ) - continue - graph = m.graph - with self._owner_lock: - self._graph_owner[graph.graph_id] = ch.channel_id - try: - self._tm.submit(graph) - except Exception as e: # noqa: BLE001 - one bad graph must not stop the node - flexkv_logger.error( - f"TE could not submit graph {graph.graph_id} from channel " - f"{ch.channel_id}: {e!r}; failing it", exc_info=True) - self._fail_graph(ch, graph.graph_id) - if had_work: - idle_spins = 0 - continue - idle_spins += 1 - if idle_spins >= 1000: - # Idle: futex wait on ctrl wake counter. - snapshot = self._ctrl.get_wake() if self._ctrl else 0 - # Re-check after snapshot — necessary to avoid lost wakeup. - any_pending = any( - ch._submit_r.value != ch._submit_w.value - for ch in self._channels - ) - if any_pending: - idle_spins = 0 - continue - if self._ctrl is not None: - self._ctrl.wait(snapshot, - timeout_ns=int(0.1 * 1_000_000_000)) - idle_spins = 0 - - def _poll_results(self) -> None: - while not self._stop.is_set(): - try: - completed = self._tm.wait(timeout=0.05) - except Exception as e: # pragma: no cover - flexkv_logger.error(f"TE result poll error: {e}") - time.sleep(0.01) - continue - if not completed: - continue - # Group completed ops by owner channel. - by_channel: Dict[int, List[CompletedOp]] = {} - for op in completed: - with self._owner_lock: - owner = self._graph_owner.get(op.graph_id) - if op.op_id == -1: - # Terminal message (completed OR failed) — drop the - # mapping after we've grouped, or failed graphs leak it. - self._graph_owner.pop(op.graph_id, None) - if owner is None: - flexkv_logger.warning( - f"TE got completed op for unknown graph {op.graph_id}" - ) - continue - by_channel.setdefault(owner, []).append(op) - for ch_id, ops in by_channel.items(): - if not (0 <= ch_id < len(self._channels)): - continue - try: - self._channels[ch_id].result_send(ops) - except Exception as e: # noqa: BLE001 - keep serving the other channels - flexkv_logger.error( - f"TE could not deliver {len(ops)} completion(s) to channel " - f"{ch_id} (graphs {sorted({op.graph_id for op in ops})}): {e!r}", - exc_info=True) - - def _fail_graph(self, ch, graph_id: int) -> None: - """Tell the owning CE that `graph_id` is over and failed when the TE could - not run it, so its task errors out instead of waiting forever.""" - from flexkv.common.transfer import CompletedOp - with self._owner_lock: - self._graph_owner.pop(graph_id, None) - try: - ch.result_send([CompletedOp(graph_id=graph_id, op_id=-1, transfer_type=None, - num_blocks=0, num_bytes=0, wait_ms=0.0, - xfer_ms=0.0, e2e_ms=0.0, failed=True)]) - except Exception as e: # noqa: BLE001 - flexkv_logger.error( - f"TE could not report the failure of graph {graph_id} to channel " - f"{ch.channel_id}: {e!r}") - - -def te_shm_main(model_config: ModelConfig, - cache_config: CacheConfig, - gpu_register_port: str, - server_id: str, - num_channels: int, - start_event, - ready_event, - stop_event) -> None: - """Entrypoint for the TE subprocess in `mode="shm"`. - - Mirrors `TransferManagerInterProcessHandle._process_worker` but replaces - the single mp.Pipe with N shm channels. - - Critical ordering: shm channel files (`flexkv_te_ctrl_*`, - `flexkv_te_ch_*_*`) must exist before any CE attaches. We therefore - create the dispatcher's channels FIRST (so CEs can open the files), then - bring up the TransferManager (which blocks on GPU registration), then - bind the TM into the dispatcher and flip the ready flag. - """ - from flexkv.transfer_manager import TransferManager - dispatcher = None - tm = None - try: - os.environ["MPI4PY_RC_INITIALIZE"] = "false" - - # Phase 1: create shm channels — CE side can attach now. - dispatcher = _TEShmDispatcher(server_id, num_channels) - dispatcher.setup_channels() - # Signal start (but not ready) so the parent's `_start_event.wait()` - # returns. Ready flag is set later by start_dispatch(). - start_event.set() - - # Phase 2: build and start the TransferManager. This blocks waiting - # for GPU clients to register over the zmq gpu_register_port. - tm = TransferManager(model_config, cache_config, gpu_register_port) - tm.initialize_transfer_engine() - tm.start() - # TransferManager binds a GPU-control REP socket in __init__; every - # deployment mode must service it or a client suspend/resume call - # stalls for the full 120s RCVTIMEO. - tm.start_gpu_control_listener() - - # Phase 3: bind TM, flip ready flag, launch poll threads. - dispatcher.start_dispatch(tm) - ready_event.set() - - # Block until parent terminates us. - while not stop_event.is_set(): - stop_event.wait(timeout=1.0) - except Exception as e: - flexkv_logger.error(f"te_shm_main failed: {e}", exc_info=True) - finally: - if dispatcher is not None: - try: - dispatcher.shutdown() - except Exception: - pass - if tm is not None: - try: - tm.shutdown() - except Exception: - pass diff --git a/flexkv/transfer_manager.py b/flexkv/transfer_manager.py index 62bf21b32..f312ed9c1 100644 --- a/flexkv/transfer_manager.py +++ b/flexkv/transfer_manager.py @@ -1557,107 +1557,6 @@ def shutdown(self) -> None: flexkv_logger.info("TransferManagerMultiNodeHandle shutdown complete") -def shm_te_clears_cuda_visible_devices(cvd: Optional[str], gpus_needed: int) -> bool: - """Whether TransferManagerShmTEProcess drops CUDA_VISIBLE_DEVICES for the TE. - - The TE opens every registered GPU buffer with cudaIpcOpenMemHandle, so it - must number the devices exactly as the workers did when they registered: - - * one CUDA_VISIBLE_DEVICES for the whole engine (sglang, single-GPU vLLM): - the workers register logical ids inside that namespace, so the TE has to - INHERIT it. Clearing it renumbers the devices and puts the TE on the wrong - physical GPUs (or, on one GPU, fails with "device >= 0 && device < num_gpus"). - * one CUDA_VISIBLE_DEVICES per DP rank (vLLM DP, each rank pinned to its own - device): the parent sees fewer GPUs than this node's TE has to reach, and - the ids the workers registered are physical. Only then is it cleared. - """ - if cvd is None: - return False - visible = len([d for d in cvd.split(",") if d.strip()]) - return visible < gpus_needed - - -class TransferManagerShmTEProcess: - """Spawns the single TE subprocess for the multi-DP shm path. - - The bootstrap (instance 0, dp 0) creates this; CEs in other DP processes - just connect via `TransferManagerHandle(mode="shm", shm_server_id=..., - shm_channel_id=...)`. - """ - - def __init__(self, - model_config: ModelConfig, - cache_config: CacheConfig, - gpu_register_port: str, - server_id: str, - num_channels: int): - self.model_config = model_config - self.cache_config = cache_config - self.gpu_register_port = gpu_register_port - self.server_id = server_id - self.num_channels = num_channels - - self.mp_ctx = mp.get_context("spawn") - self._start_event = self.mp_ctx.Event() - self._ready_event = self.mp_ctx.Event() - self._stop_event = self.mp_ctx.Event() - self.process: Optional[Process] = None - - def start(self) -> None: - if self.process is not None and self.process.is_alive(): - return - from flexkv.transfer.shm_channel_handle import te_shm_main - # mp.Process(spawn) inherits the parent's env. The TE keeps the parent's - # CUDA_VISIBLE_DEVICES whenever that namespace covers the GPUs it serves - # (the workers registered ids inside it); it is cleared only when the - # parent is pinned to fewer GPUs than the node's TE must reach. See - # shm_te_clears_cuda_visible_devices. Pop / restore around .start(). - cvd = os.environ.get("CUDA_VISIBLE_DEVICES") - gpus_needed = self.model_config.instance_num * self.model_config.gpus_per_node - clear_cvd = shm_te_clears_cuda_visible_devices(cvd, gpus_needed) - if cvd is not None: - flexkv_logger.info( - f"TransferManagerShmTEProcess: parent CUDA_VISIBLE_DEVICES={cvd!r}, " - f"TE serves {gpus_needed} GPU(s) on this node -> " - f"{'clearing it for the TE (per-rank pinned layout)' if clear_cvd else 'the TE inherits it'}") - _saved_cuda = os.environ.pop("CUDA_VISIBLE_DEVICES", None) if clear_cvd else None - try: - self.process = self.mp_ctx.Process( - target=te_shm_main, - args=(self.model_config, - self.cache_config, - self.gpu_register_port, - self.server_id, - self.num_channels, - self._start_event, - self._ready_event, - self._stop_event), - daemon=False, - ) - self.process.start() - finally: - if _saved_cuda is not None: - os.environ["CUDA_VISIBLE_DEVICES"] = _saved_cuda - self._start_event.wait() - flexkv_logger.info( - f"TransferManagerShmTEProcess started, PID={self.process.pid}, " - f"server_id={self.server_id}, channels={self.num_channels}" - ) - - def is_ready(self) -> bool: - return self._ready_event.is_set() - - def shutdown(self, timeout: float = 5.0) -> None: - if self.process is None: - return - self._stop_event.set() - self.process.join(timeout=timeout) - if self.process.is_alive(): - self.process.terminate() - self.process.join() - self.process = None - - class TransferManagerHandle: def __init__(self, model_config: ModelConfig, @@ -1686,22 +1585,8 @@ def __init__(self, self._handle: TransferManagerHandleBase = TransferManagerMultiNodeHandle( model_config, cache_config, gpu_register_port, master_host, master_ports ) - elif mode == "shm": - # Multi-DP path: each CE process gets a dedicated ShmChannel to a - # single TE subprocess. The TE is created by the bootstrap process - # via TransferManagerShmTEProcess; clients attach by server_id. - from flexkv.transfer.shm_channel_handle import ( - TransferManagerShmChannelHandle, - ) - server_id = kwargs["shm_server_id"] - channel_id = kwargs["shm_channel_id"] - self._handle: TransferManagerHandleBase = TransferManagerShmChannelHandle( - model_config, cache_config, server_id, channel_id - ) else: - raise ValueError( - f"Invalid mode: {mode}, must be process, thread, remote, or shm" - ) + raise ValueError(f"Invalid mode: {mode}, must be process, thread or remote") def start(self) -> None: self._handle.start() diff --git a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py index 1d382f26f..e35ac7bcb 100644 --- a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -113,8 +113,7 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, # addressing it as device 0. os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id) node_name = _node_name(rank) - # FlexKV's own IPC names are per node too (its radix regions get the - # node name appended through the same override). + # FlexKV's own IPC names are per node too. recv_port = f"ipc:///tmp/flexkv_{cluster_id}_{node_name}" os.environ.update({ "FLEXKV_ENABLE_RADIXSHMEM": "1", @@ -142,8 +141,8 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, tokens_per_block=TOKENS_PER_BLOCK, enable_cpu=True, enable_ssd=False, enable_remote=False, num_cpu_blocks=NUM_CPU_BLOCKS, - # Peer reuse follows the radixshmem YAML (expected_min_nodes=2 below), - # not enable_p2p_cpu. + # Peer reuse follows the radix-server (started with cluster flags + # below), not enable_p2p_cpu. ) report = {"rank": rank} diff --git a/tests/radixshmem/test_e2e_radix_shmem.py b/tests/radixshmem/test_e2e_radix_shmem.py index 70074f27f..7418857e3 100644 --- a/tests/radixshmem/test_e2e_radix_shmem.py +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -1,16 +1,17 @@ """End-to-end test of FLEXKV_ENABLE_RADIXSHMEM=1 on one node: one or two DP scheduler -processes, one radix-server, one shared transfer engine. +processes and one radix-server, on FlexKV's own process model (dp_size 1: the +KVTaskEngine and its TE subprocess in the engine process; dp_size 2: dp0 embeds +the KVServer whose KVTaskEngine and TE serve both DPs). For every ``dp_size`` in the parametrization: * the test starts the operator's radix-server (index + SlotStore, the - node's CPU KV pool) with nothing but a name and a byte budget; dp0's - KVManager hands it FlexKV's geometry, adopts the slot counts it plans and - spawns the single TE; every other DP attaches to both by name and feeds - the TE over its own shm channel with a disjoint graph/op id range. Which - server to attach to is the ``server.name`` of a small YAML written per run + node's CPU KV pool) with nothing but a name and a byte budget; every + KVManager hands it FlexKV's geometry and adopts the slot counts it plans, + and so does the process that builds the KVTaskEngine. Which server to + attach to is the ``server.name`` of a small YAML written per run (FLEXKV_RADIXSHMEM_CONFIG_PATH), which is how a deployment names it too. - * Phase 1: every DP PUTs its own requests concurrently through the shared TE. + * Phase 1: every DP PUTs its own requests concurrently. * Phase 2 (dp_size > 1): dp0 PUTs a prefix that dp1 then finds with ``get_match`` -- the shared index is what the radixshmem path exists for. * Phase 3: every DP writes a rank-specific byte pattern into its GPU blocks, @@ -68,8 +69,8 @@ def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, barrier, result_q) -> None: """Full lifecycle of one DP scheduler process.""" # Before the first flexkv import: GLOBAL_CONFIG_FROM_ENV is read at import. - # All DP procs share one TE, so they must agree on server_recv_port (and - # therefore on the gpu_register_port the TE listens on). + # With dp_size > 1 the DPs are KVServer clients and must agree on + # server_recv_port (and therefore on the gpu_register_port the TE listens on). recv_port = f"ipc:///tmp/flexkv_{server_id}" os.environ.update({ "FLEXKV_ENABLE_RADIXSHMEM": "1", @@ -103,9 +104,8 @@ def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, try: kvm = KVManager(model_config, cache_config, dp_client_id=dp_client_id) kvm.start() - # Each DP drives its own GPU (device id = dp id; the TE opens every - # DP's IPC handles because total_gpus > 1 clears CUDA_VISIBLE_DEVICES - # for it when dp_size > 1). + # Each DP drives its own GPU (device id = dp id); no process pins + # CUDA_VISIBLE_DEVICES, so the TE sees every device the DPs registered. tp_proc, gpu_tensors = start_tp_client( kvm, dp_client_id, dp_client_id, model_config, cache_config, NUM_GPU_BLOCKS) wait_kv_manager_ready(kvm, timeout=120) diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index ccaf904e7..b40c8a399 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -748,7 +748,7 @@ def test_radix_config_defaults(): cfg = load_radixshmem_config(None) assert cfg.path is None assert cfg.server_name == "/flexkv" and cfg.endpoint == "" and cfg.ready_timeout_s == 600.0 - assert cfg.te_server_id == "flexkv" + assert cfg.default_endpoint == "unix:///dev/shm/flexkv.sock" assert cfg.client.prefetch_timeout_ms == 5000 and cfg.client.prefetch_max_inflight == 128 assert cfg.client.max_outstanding == 256 assert "radix-server /flexkv" in cfg.describe() @@ -766,7 +766,8 @@ def test_radix_config_file(tmp_path): prefetch_max_inflight: 300 """) cfg = load_radixshmem_config(path) - assert cfg.path == path and cfg.server_name == "/prod/kv" and cfg.te_server_id == "prod_kv" + assert cfg.path == path and cfg.server_name == "/prod/kv" + assert cfg.default_endpoint == "unix:///dev/shm/prod_kv.sock" assert cfg.endpoint == "10.0.0.2:7000" and cfg.ready_timeout_s == 900.0 assert cfg.client.prefetch_timeout_ms == 1000 and cfg.client.max_outstanding == 512 assert cfg.client.prefetch_max_inflight == 300 @@ -808,7 +809,7 @@ def test_radix_config_env_singleton_reloads_on_change(tmp_path, monkeypatch): set_radixshmem_config(_radix_config(name="/pinned")) assert get_radixshmem_config().server_name == "/pinned" set_radixshmem_config(None) - assert get_radixshmem_config().te_server_id == "other" + assert get_radixshmem_config().default_endpoint == "unix:///dev/shm/other.sock" # ============================================================================= diff --git a/tests/test_shm_channel.py b/tests/test_shm_channel.py deleted file mode 100644 index d21577b06..000000000 --- a/tests/test_shm_channel.py +++ /dev/null @@ -1,333 +0,0 @@ -"""Round-trip tests for the CE↔TE ShmChannel. - -Run with:: - - python3 -m pytest tests/test_shm_channel.py -v - -The test forks N producer processes and one consumer; each producer sends K -graph-shaped messages and reads back acks via its own result ring. -""" -from __future__ import annotations - -import multiprocessing as mp -import os -import pickle -import threading -import time - -import numpy as np -import pytest - -from flexkv.common.transfer import CompletedOp -from flexkv.transfer.shm_channel import ShmChannel, ShmControlBlock - - -SERVER_ID = "shm_channel_test" - - -def _producer(channel_id: int, n: int, server_id: str) -> None: - ch = ShmChannel(server_id, channel_id, create=False) - for i in range(n): - ch.submit_send({"channel": channel_id, "seq": i, "data": b"x" * 1024}) - # Wait for echoed acks (CompletedOps carrying channel/seq in graph_id/op_id). - received = 0 - while received < n: - msgs = ch.result_recv(timeout_s=2.0) - received += len(msgs) - for m in msgs: - assert m.graph_id == channel_id - - -def _consumer(num_channels: int, total_per_channel: int, server_id: str) -> None: - ctrl = ShmControlBlock(server_id, create=True) - channels = [ - ShmChannel(server_id, i, create=True) for i in range(num_channels) - ] - ctrl.set_ready() - pending = num_channels * total_per_channel - while pending > 0: - had_work = False - for ch in channels: - msgs = ch.submit_recv() - if msgs: - had_work = True - for m in msgs: - # Echo back the (channel, seq) as a CompletedOp — the result - # ring is now typed to CompletedOp records. - ch.result_send([CompletedOp(graph_id=m["channel"], - op_id=m["seq"])]) - pending -= 1 - if not had_work: - time.sleep(0.001) - for ch in channels: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def _cleanup_shm(server_id: str, num_channels: int) -> None: - for ch_id in range(num_channels): - path = f"/dev/shm/flexkv_te_ch_{server_id}_{ch_id}" - if os.path.exists(path): - os.unlink(path) - ctrl_path = f"/dev/shm/flexkv_te_ctrl_{server_id}" - if os.path.exists(ctrl_path): - os.unlink(ctrl_path) - - -def test_n_producer_one_consumer(): - server_id = f"{SERVER_ID}_npm" - num_channels = 4 - total_per_channel = 32 - - _cleanup_shm(server_id, num_channels) - - ctx = mp.get_context("spawn") - consumer = ctx.Process( - target=_consumer, - args=(num_channels, total_per_channel, server_id), - ) - consumer.start() - - # Wait for consumer to set up shm files. - deadline = time.monotonic() + 5.0 - while time.monotonic() < deadline: - if os.path.exists(f"/dev/shm/flexkv_te_ctrl_{server_id}"): - break - time.sleep(0.01) - else: - consumer.terminate() - pytest.fail("consumer never created control block") - - # Wait for ready flag via control block. - ctrl = ShmControlBlock(server_id, create=False) - assert ctrl.wait_ready(timeout_s=5.0) - ctrl.close() - - producers = [ - ctx.Process( - target=_producer, args=(i, total_per_channel, server_id) - ) - for i in range(num_channels) - ] - for p in producers: - p.start() - for p in producers: - p.join(timeout=10.0) - assert p.exitcode == 0, f"producer {p.pid} exit code {p.exitcode}" - - consumer.join(timeout=10.0) - assert consumer.exitcode == 0 - - -def test_single_round_trip_local(): - server_id = f"{SERVER_ID}_local" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - ch = ShmChannel(server_id, 0, create=True) - try: - ch.submit_send("hello") - ch.submit_send({"k": 42}) - msgs = ch.submit_recv() - assert msgs == ["hello", {"k": 42}] - - # Result ring carries fixed-width CompletedOp records; all fields must - # round-trip, including the transfer_type string, the worker-measured - # durations and the -1 sentinel. - sent = [ - CompletedOp(graph_id=7, op_id=3, transfer_type="H2D", - num_blocks=12, num_bytes=98304, - wait_ms=0.25, xfer_ms=1.5, e2e_ms=2.125), - CompletedOp(graph_id=7, op_id=-1), # graph-completed sentinel - ] - ch.result_send(sent) - out = ch.result_recv(timeout_s=0.0) - assert out == sent - assert (out[0].wait_ms, out[0].xfer_ms, out[0].e2e_ms) == (0.25, 1.5, 2.125) - assert out[1].is_graph_completed() - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def _explode(): - raise ValueError("poisoned record") - - -class _Poison: - """Pickles fine, fails to unpickle: what a corrupted submit slot looks like.""" - - def __reduce__(self): - return (_explode, ()) - - -def test_undecodable_submit_record_is_dropped_not_fatal(): - """A record the TE cannot decode is dropped with an error and the ring keeps - moving; re-raising on every poll would stall every channel on the node.""" - server_id = f"{SERVER_ID}_poison{os.getpid()}" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - ch = ShmChannel(server_id, 0, create=True) - try: - ch.submit_send(_Poison()) - ch.submit_send("after the poison") - assert ch.submit_recv() == ["after the poison"] - assert ch.submit_recv() == [] # the ring advanced past both - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def test_every_transfer_type_has_a_wire_index(): - """The TE sends `op.transfer_type.value` for every non-VIRTUAL op, so each - TransferType member must round-trip. LAYERWISE was missing: its KeyError in - the TE's result thread stopped every completion on the node.""" - from flexkv.common.transfer import TransferType - from flexkv.transfer.shm_channel import decode_completed_op, encode_completed_op - for tt in TransferType: - name = None if tt == TransferType.VIRTUAL else tt.value - op = CompletedOp(graph_id=1, op_id=2, transfer_type=name, num_blocks=1, num_bytes=8) - back = decode_completed_op(memoryview(encode_completed_op(op)), 0) - assert back == op, tt - # the member itself is accepted too - op = CompletedOp(graph_id=1, op_id=2, transfer_type=TransferType.LAYERWISE, num_blocks=1) - assert decode_completed_op(memoryview(encode_completed_op(op)), 0).transfer_type == "LAYERWISE" - - -def test_unknown_transfer_type_is_sent_as_none(): - """A name outside the table degrades to None instead of raising.""" - from flexkv.transfer.shm_channel import decode_completed_op, encode_completed_op - op = CompletedOp(graph_id=1, op_id=2, transfer_type="NOT_A_TRANSFER_TYPE", num_blocks=1) - back = decode_completed_op(memoryview(encode_completed_op(op)), 0) - assert back.transfer_type is None and (back.graph_id, back.op_id, back.num_blocks) == (1, 2, 1) - - -def test_result_record_carries_durations_and_fits_a_slot(): - """The fixed-width record must hold the #297 durations (f64) and still fit - the default 64 B result slot; a failed op keeps its flag alongside them.""" - from flexkv.transfer.shm_channel import ( - COMPLETED_OP_WIRE_SIZE, DEFAULT_RESULT_SLOT_SIZE, - decode_completed_op, encode_completed_op) - assert COMPLETED_OP_WIRE_SIZE <= DEFAULT_RESULT_SLOT_SIZE - op = CompletedOp(graph_id=1 << 40, op_id=(1 << 40) + 5, transfer_type="D2H", - num_blocks=3, num_bytes=3 * 4096, - wait_ms=12.75, xfer_ms=0.001, e2e_ms=1e6, failed=True) - rec = encode_completed_op(op) - assert len(rec) == COMPLETED_OP_WIRE_SIZE - back = decode_completed_op(memoryview(rec), 0) - assert back == op - assert back.is_graph_failed() is False and back.failed - - -def test_submit_fragmentation(): - """A payload larger than one slot must fragment and round-trip intact, - interleaved with small single-slot messages.""" - server_id = f"{SERVER_ID}_frag" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - # Small slots so a modest payload spans many fragments. - ch = ShmChannel(server_id, 0, create=True, - submit_slots=1024, slot_size=4096) - try: - big = {"arr": np.arange(200_000, dtype=np.int64)} # ~1.5 MB > slot - small = {"k": 1} - ch.submit_send(small) - ch.submit_send(big) - ch.submit_send(small) - msgs = ch.submit_recv() - assert len(msgs) == 3 - assert msgs[0] == small - assert np.array_equal(msgs[1]["arr"], big["arr"]) - assert msgs[2] == small - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def test_submit_payload_too_large(): - """A payload that can't fit the whole ring is rejected, not deadlocked.""" - server_id = f"{SERVER_ID}_toobig" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - ch = ShmChannel(server_id, 0, create=True, - submit_slots=8, slot_size=4096) - try: - with pytest.raises(ValueError): - ch.submit_send(b"x" * (8 * 4096)) # needs more fragments than slots - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def _make_big_graph(nbytes: int) -> dict: - """A graph-shaped payload that pickles to ~nbytes (block-id arrays dominate).""" - n = nbytes // 16 # two int64 arrays - return {"src": np.arange(n, dtype=np.int64), - "dst": np.arange(n, dtype=np.int64)} - - -def test_submit_big_graph_500kb(): - """A ~500 KB graph fragments across the default 32 KB slots and round-trips.""" - server_id = f"{SERVER_ID}_big500" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - ch = ShmChannel(server_id, 0, create=True) # default 32 KB / 8192 - try: - big = _make_big_graph(500 * 1024) - blob_sz = len(pickle.dumps(big, protocol=pickle.HIGHEST_PROTOCOL)) - assert blob_sz > 500 * 1024, f"payload only {blob_sz} B" - assert blob_sz > ch.slot_size, "payload must exceed one slot" - ch.submit_send(big) - out = ch.submit_recv() - assert len(out) == 1 - assert np.array_equal(out[0]["src"], big["src"]) - assert np.array_equal(out[0]["dst"], big["dst"]) - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() - - -def test_submit_small_graphs_high_rate(): - """4096 small multi-slot graphs at ~4096/s: a producer thread submits while a - consumer thread drains, verifying no loss, correct order, and no ring-full - stall at the default 8192-slot capacity.""" - server_id = f"{SERVER_ID}_hirate" - _cleanup_shm(server_id, 1) - ctrl = ShmControlBlock(server_id, create=True) - ch = ShmChannel(server_id, 0, create=True) # default 32 KB / 8192 - n_msgs = 4096 - small = {"payload": b"x" * (2 * ch.slot_size)} # spans ~3 slots each - received: list = [] - - def consumer() -> None: - while len(received) < n_msgs: - received.extend(ch.submit_recv()) - - try: - t = threading.Thread(target=consumer, daemon=True) - t.start() - start = time.monotonic() - for i in range(n_msgs): - ch.submit_send({"seq": i, **small}) - t.join(timeout=30.0) - elapsed = time.monotonic() - start - assert len(received) == n_msgs, f"got {len(received)}/{n_msgs}" - assert [m["seq"] for m in received] == list(range(n_msgs)), "order/loss" - rate = n_msgs / elapsed - assert rate >= 4096, f"throughput {rate:.0f}/s below 4096/s target" - finally: - ch.close() - ch.unlink() - ctrl.close() - ctrl.unlink() diff --git a/tests/test_shm_te_cvd.py b/tests/test_shm_te_cvd.py deleted file mode 100644 index ec5701e31..000000000 --- a/tests/test_shm_te_cvd.py +++ /dev/null @@ -1,17 +0,0 @@ -"""CUDA_VISIBLE_DEVICES policy of the radixshmem shm TE subprocess.""" -import pytest - -from flexkv.transfer_manager import shm_te_clears_cuda_visible_devices - - -@pytest.mark.parametrize("cvd, needed, clears", [ - (None, 2, False), # nothing set: the TE sees every GPU, ids are physical - ("2,3", 2, False), # sglang: one namespace for all TP workers -> inherit - ("0,1,2,3", 2, False), # a wider namespace than needed still covers the TE - ("5", 1, False), # single-GPU deployment on a restricted device -> inherit - ("3", 8, True), # vLLM DP: rank pinned to one device, TE serves 8 -> clear - ("", 2, True), # empty value hides every GPU; drop it - ("GPU-aaaa,GPU-bbbb", 2, False), # UUID form counts the same way -]) -def test_shm_te_cvd_policy(cvd, needed, clears): - assert shm_te_clears_cuda_visible_devices(cvd, needed) is clears From d1d783651724440ae4421e5d67f8101de39821d5 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 15:04:28 +0800 Subject: [PATCH 14/21] radixshmem: configure the server by name (FLEXKV_RADIXSHMEM_SERVER_NAME); fixed attach defaults FlexKV's side of radixshmem mode is one setting: FLEXKV_RADIXSHMEM_SERVER_NAME (default /flexkv), the --name of the node's radix-server (shm_radix_bootstrap.radix_server_name(); a shm name: starts with '/', no whitespace). The gRPC endpoint is the socket radixshmem derives from the name (unix:///dev/shm/.sock). The ready wait (READY_TIMEOUT_S = 600 s: a late server start, the SlotStore prefault, a cluster rendezvous), the prefetch pull deadline (PREFETCH_TIMEOUT_MS = 5000), the peer pulls in flight per KVTaskEngine (PREFETCH_MAX_INFLIGHT = 128) and the RadixClient job cap (MAX_OUTSTANDING = 256) are constants in flexkv/server/shm_radix_bootstrap.py. attach_radix_client, adopt_radix_server and radix_server_is_distributed take the server name; the planner reads the constants. flexkv/common/radixshmem_config.py (the YAML behind FLEXKV_RADIXSHMEM_CONFIG_PATH) is gone, and so is examples/radixshmem_configs/radixshmem.yaml; the launch scripts live in examples/radixshmem/. Tests: the YAML tests are replaced by tests of the variable, the e2e tests pass the server name through the environment, and the RDMA cluster test gives its two servers distinct names (hence distinct sockets) instead of an endpoint override. docs/radixshmem/config_zh.md and the CHANGELOG describe the variable and the fixed defaults. Verified: 442 CPU tests (the CI unit files, the radixshmem suite, the sglang store protocol); the GPU e2e tests/radixshmem/test_e2e_radix_shmem.py passes for dp_size 1 and 2. --- CHANGELOG.md | 11 +- docs/radixshmem/config_zh.md | 198 +++++---------- .../radix_server_multi_node.sh | 0 .../radix_server_single_node.sh | 0 examples/radixshmem_configs/radixshmem.yaml | 19 -- flexkv/cache/radix_shmem_planner.py | 18 +- flexkv/common/config.py | 13 +- flexkv/common/radixshmem_config.py | 240 ------------------ flexkv/server/shm_radix_bootstrap.py | 64 +++-- tests/radixshmem/radix_e2e_common.py | 10 - .../radixshmem/test_e2e_radix_prefetch_p2p.py | 21 +- tests/radixshmem/test_e2e_radix_shmem.py | 14 +- tests/radixshmem/test_radix_shmem_engine.py | 147 ++++------- 13 files changed, 181 insertions(+), 574 deletions(-) rename examples/{radixshmem_configs => radixshmem}/radix_server_multi_node.sh (100%) rename examples/{radixshmem_configs => radixshmem}/radix_server_single_node.sh (100%) delete mode 100644 examples/radixshmem_configs/radixshmem.yaml delete mode 100644 flexkv/common/radixshmem_config.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ecaadd235..5875d3e55 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,14 +10,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Feature Universal: -- radixshmem mode: the CPU tier is an operator-run `radix-server` (the radixshmem project; requires radixshmem f939910 or later). FlexKV no longer creates a server: the operator starts one per node with a name and a byte budget (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`), FlexKV's clients hand it the geometry (`shmradix.RadixClient(name, Geometry)`: tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), the server plans the slot counts from its budget and every FlexKV process adopts them into `CacheConfig.num_cpu_blocks` / `swa.num_slots` (`shm_radix_bootstrap.adopt_geometry`; `cpu_cache_gb` no longer sizes the CPU tier in this mode). A server serving another geometry is refused (`GeometryMismatch`), so every engine on one server runs the same model, page size and SWA configuration. `RadixServerProcess`, `build_radix_server_config`, `FLEXKV_RADIX_SERVER_LAUNCH_MODE`, `FLEXKV_RADIX_NODE_NAME`, `FLEXKV_RADIX_RPC_ADDRESS`, the `FLEXKV_RADIX_*` cluster variables and `FLEXKV_SHM_RADIX_ID` are gone. -- radixshmem mode is configured by one small YAML (`FLEXKV_RADIXSHMEM_CONFIG_PATH`, `flexkv/common/radixshmem_config.py`): `server` (`name`, `endpoint`, `ready_timeout_s`: which radix-server to attach to and how long to wait for it to be reachable and ready) and `client` (`prefetch_timeout_ms`, `prefetch_max_inflight`, `max_outstanding`). No file means `radix-server --name /flexkv` on the local socket. The former `cluster` / `data` / `index` sections are rejected: those settings are radix-server command-line flags now (migration table in `docs/radixshmem/config_zh.md`, launch scripts and a YAML in `examples/radixshmem_configs/`). FlexKV's own TE channel names derive from `server.name`. -- radixshmem mode brings no RHT registration chunk of its own any more (the former `REGISTER_CHUNK_TOKENS = 4096` constant, handed to the server as `index.register_chunk_size` in blocks): the chunk is the radix-server's `--register-chunk-tokens` (radixshmem's default 4096 tokens). `RadixGeometry.register_chunk_tokens` forwards a pinned value (0 = the server's), `adopt_geometry` and `CacheEngineRadixShmem.register_chunk_tokens` / `register_chunk_blocks` take the published value over, converted to blocks by radixshmem's rule (`tokens // tokens_per_block`, at least 1); `check_geometry` compares a pinned value and warns when the server's chunk is not a whole number of FlexKV blocks. -- radixshmem mode uses radixshmem's data plane: the CPU KV pool is the radix-server's SlotStore (one slot per block, attached by name in the TE and every transfer worker) and cross-node reuse is `RadixClient.pull_async` from the prefetch path (server-side RDMA READ), replacing FlexKV's own CPU allocation, `PEER2CPUTransferWorker`, mooncake wrapper and Redis address book on this path. Peer reuse follows the radix-server itself (`world_size > 1`, asked through `shm_radix_bootstrap.radix_server_is_distributed`) and no longer reads `enable_p2p_cpu`, which must stay off. radixshmem mode is CPU-tier only (`ssd_cache_gb` must be 0). Reference: `docs/radixshmem/config_zh.md`. -- radixshmem mode keeps FlexKV's process model: with one DP the KVTaskEngine runs in the engine process with its own TE subprocess, otherwise (dp_size > 1, `FLEXKV_INSTANCE_NUM` engines on one node) the DPs are clients of the node's KVServer whose KVTaskEngine and TE attach the radix-server like any other. Wherever a KVTaskEngine is built, `shm_radix_bootstrap.adopt_radix_server` first hands the server FlexKV's geometry and takes its slot counts and cluster rank over into `CacheConfig`. The embedded KVServer child now inherits `PYTHONPATH` / `LD_LIBRARY_PATH` / `PATH` along with the `FLEXKV_*` variables, so it imports what its parent imports (shmradix included) when the packages are not in site-packages. -- radixshmem planning moved out of `GlobalCacheEngine` into the subclass `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, selected by `KVTaskEngine` when `FLEXKV_ENABLE_RADIXSHMEM=1`). `GlobalCacheEngine` keeps two hooks only (`_prepare_request`, `_build_cpu_cache_engine`); its plan dataclasses, `TransferPlanHandle` and completion callbacks are back to their non-radixshmem shape. The subclass's handles roll back a plan cancelled before launch (match pin released, staged PUT slots returned), which the old planners did not, and a PUT whose planning fails after its slots were taken returns them and drops the pin before re-raising. `CacheEngineRadixShmem.take` clamps a request to the pool's size instead of letting radixshmem refuse it. -- `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`) is trimmed to what `RadixShmemCacheEngine` uses: the `CacheEngineAccel`-compatibility parameters (`device_type`, `evict_ratio`, `evict_start_threshold`, `hit_reward_seconds`, `eviction_policy`, `protected_threshold`, `tokens_per_block=-1`), `take(strict=)`, `match(gpu_matched_blocks=)`, the `mempool` view, `start()`, `store` / `cluster_rank` and the `FLEXKV_TRACE_RADIX_PEER` variable are gone (prefetch logs at debug level; the planner reports mempool metrics itself). -- `gen_hashes` and `Hasher.update` hash numpy buffers directly (`c_ext.gen_hashes_numpy` / `update_numpy`) instead of going through `torch.from_numpy`, which is not safe to call concurrently. Hashes are unchanged for int64 tokens; `gen_hashes` keeps requiring int64 and the binding now checks dtype, contiguity and sizes instead of reinterpreting the buffer. +- Add radixshmem mode (`FLEXKV_ENABLE_RADIXSHMEM=1`; requires radixshmem f939910 or later): the CPU tier is a per-node `radix-server` run by the operator (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`) that owns the radix index, the SlotStore and the cross-node RDMA pull. FlexKV attaches with `shmradix.RadixClient(name, Geometry)`: it hands the server its slot geometry (tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), adopts the slot counts the server plans from its budget into `CacheConfig` (`shm_radix_bootstrap.adopt_radix_server`; `cpu_cache_gb` has no effect in this mode) and uses the server's SlotStore as the CPU pool in the TE and every transfer worker; cross-node reuse is `RadixClient.pull_async` from the prefetch path whenever the server runs as a cluster. `FLEXKV_RADIXSHMEM_SERVER_NAME` (default `/flexkv`) names the server to attach to; the endpoint is the one radixshmem derives from the name, and the ready wait (600 s) and the prefetch limits are fixed defaults. Process model as in the other modes: engine mode for one DP, the node's KVServer for dp_size > 1 or several `FLEXKV_INSTANCE_NUM` engines. Planning is `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, a `GlobalCacheEngine` subclass over the `_prepare_request` / `_build_cpu_cache_engine` hooks) on `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`); a plan cancelled before launch or failing during planning returns its slots and match pin. CPU tier only (`enable_ssd`, `enable_remote`, `enable_p2p_*` must be off). Reference: `docs/radixshmem/config_zh.md`, launch scripts in `examples/radixshmem/`. +- `KVServer.create_server(inherit_env=False)` passes `PYTHONPATH`, `LD_LIBRARY_PATH` and `PATH` to the server child along with the `FLEXKV_*` variables. +- `gen_hashes` / `Hasher.update` hash numpy buffers directly (`c_ext.gen_hashes_numpy` / `update_numpy`, safe to call concurrently); `gen_hashes_numpy` validates dtype (int64 tokens, uint64 hashes), contiguity and sizes. Targeting SGLang: - The native FlexKV backend is available in upstream SGLang `v0.5.16` and later; no patch is required ([sglang#29701](https://github.com/sgl-project/sglang/pull/29701)) diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 73cc6b46d..8dca2491d 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -1,80 +1,45 @@ -# radixshmem 模式配置参考 +# radixshmem 模式配置 -FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)时,配置分两处: +FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)。配置分两处: | 谁 | 载体 | 内容 | |---|---|---| -| 运维 | `radix-server` 命令行,每节点一个进程 | 名字、SlotStore 字节预算及 SWA 占比、hugepage、传输引擎、集群成员(etcd、网卡、rank)、索引调优 | -| FlexKV | 环境变量 + 一个很小的 YAML(`FLEXKV_RADIXSHMEM_CONFIG_PATH`) | 是否启用、attach 哪个 server、等待多久、prefetch 限额 | -| FlexKV 推导 | `ModelConfig` / `CacheConfig` | 几何:每 block 的 token 数、一个 CPU block 的字节数、一个 SWA page 的字节数、SWA 窗口、slot 对齐 | +| 运维 | `radix-server` 命令行,每节点一个进程 | 名字、SlotStore 字节预算与 SWA 占比、hugepage、传输引擎、集群成员(etcd、网卡、rank)、索引调优 | +| FlexKV | 两个环境变量 | 是否启用、attach 哪个 server | -FlexKV **不再创建 radix-server**。它只实例化 `shmradix.RadixClient`:第一个 client 把几何交给 server, -server 按自己的字节预算规划各池的 slot 数并发布;之后每个 FlexKV 进程 attach 时把这些 slot 数采纳到 -`CacheConfig`(`num_cpu_blocks`、`swa.num_slots`)。也就是说,在这个模式下 CPU 层的容量由 -`radix-server --data-bytes` 决定,`cpu_cache_gb` 只是占位。 +FlexKV 只实例化 `shmradix.RadixClient`:第一个 client 把几何(每 block 的 token 数、一个 CPU block 和一个 SWA page +的字节数、SWA 窗口、slot 对齐)交给 server,server 按 `--data-bytes` 和 `--swa-ratio` 规划各池的 slot 数并发布, +每个 FlexKV 进程 attach 时把 slot 数采纳到 `CacheConfig`(`num_cpu_blocks`、`swa.num_slots`)。CPU 层容量由 +`radix-server --data-bytes` 决定,`cpu_cache_gb` 在该模式下不起作用。 -实现:`flexkv/common/radixshmem_config.py`(YAML)、`flexkv/server/shm_radix_bootstrap.py`(几何、attach、采纳)。 -radixshmem 侧的接口见 radixshmem 仓库 `python/README.md`。 - ---- +实现:`flexkv/server/shm_radix_bootstrap.py`(几何、attach、采纳、固定参数)。 +radixshmem 侧接口见 radixshmem 仓库 `python/README.md`。 ## 1. 环境变量 | 变量 | 默认 | 说明 | |---|---|---| -| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | 模式总开关。`1` 时 CPU 层由 radixshmem 承担。进程模型沿用 FlexKV 原有的:`dp_size=1` 且单实例时 KVTaskEngine 在引擎进程内并自带 TE 子进程;`dp_size>1` 或多实例时走 server-client 模式,每节点一个 KVServer,其中的 KVTaskEngine 和 TE attach radix-server。在 `flexkv` 首次 import 前设置。 | -| `FLEXKV_RADIXSHMEM_CONFIG_PATH` | 空 | 第 2 节 YAML 的路径。为空时全部取默认值:attach 本机 `radix-server --name /flexkv`。 | - -另有两个 FlexKV 通用变量在该模式下有约束: - -- `FLEXKV_CPU_LAYOUT` 必须是 `BLOCKFIRST`。一个 SlotStore slot 就是一个连续的 CPU block,LAYERFIRST 给不出这个布局。 -- `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID`:同一节点上多个推理引擎共享同一个 radix-server 时用来区分实例,语义与 FlexKV 原有的多实例模式相同(第 4.4 节)。 - -该模式与 `enable_ssd`、`enable_remote` 互斥,启动时报错。`enable_p2p_cpu` / `enable_p2p_ssd` 也必须为 False: -跨节点复用由 radix-server 自己完成(etcd + RDMA),在它以集群参数启动时自动开启,不经过 FlexKV 的 Redis P2P 路径。 - -已移除的变量:`FLEXKV_RADIX_SERVER_LAUNCH_MODE`(不再有嵌入式启动)、`FLEXKV_RADIX_NODE_NAME`、 -`FLEXKV_RADIX_RPC_ADDRESS`(节点身份是 `radix-server --node-name` / `--rpc-address` 的事)。 - ---- +| `FLEXKV_ENABLE_RADIXSHMEM` | `0` | `1` 启用。在 `flexkv` 首次 import 前设置。 | +| `FLEXKV_RADIXSHMEM_SERVER_NAME` | `/flexkv` | attach 的 `radix-server --name`,即索引 shm 名,以 `/` 开头。gRPC 端点是 radixshmem 由名字派生的 `unix:///dev/shm/.sock`。 | +| `FLEXKV_CPU_LAYOUT` | | 必须是 `BLOCKFIRST`:一个 SlotStore slot 就是一个连续的 CPU block。 | +| `FLEXKV_INSTANCE_NUM` / `FLEXKV_INSTANCE_ID` | `1` / `0` | 同一节点上多个推理引擎共享一个 radix-server 时区分实例(第 4.4 节)。 | -## 2. YAML 字段 +该模式只承担 CPU 层:`enable_ssd`、`enable_remote`、`enable_p2p_cpu`、`enable_p2p_ssd` 必须关闭,启动时校验。 +跨节点复用由 radix-server 完成(etcd + RDMA),server 以集群参数启动时自动开启。 -两个段,都可省略。 +进程模型与 FlexKV 其它模式一致:`dp_size=1` 且单实例时 KVTaskEngine 在引擎进程内,TE 是它的子进程; +`dp_size>1` 或多实例时每节点一个 KVServer,DP 进程是它的 client,KVServer 里的 KVTaskEngine 和 TE attach radix-server。 -### 2.1 `server`:attach 哪个 radix-server +## 2. 固定参数 -| 键 | 默认 | 说明 | -|---|---|---| -| `name` | `/flexkv` | `radix-server --name`,即索引 shm 名。以 `/` 开头。也派生默认 socket(第 5 节)。 | -| `endpoint` | 空 | gRPC 端点。空为 `unix:///dev/shm/.sock`;server 以 `--endpoint` 改成 TCP 或别的路径时这里写同一个值。 | -| `ready_timeout_s` | `600` | 一个 FlexKV 进程等 server **可达且 ready** 的总时长。覆盖运维晚起 server、SlotStore prefault、集群 rendezvous(server 的 `--bootstrap-timeout`)。超时报错并给出启动命令。 | - -### 2.2 `client`:FlexKV 侧参数 +attach 的其余参数是 `flexkv/server/shm_radix_bootstrap.py` 里的常量: -| 键 | 默认 | 说明 | +| 常量 | 值 | 含义 | |---|---|---| -| `prefetch_timeout_ms` | `5000` | 一次 prefetch 拉取的服务端超时。到期后 job 以本地命中的部分完成。 | -| `prefetch_max_inflight` | `128` | 每个 KVTaskEngine 在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。需小于 `max_outstanding`。 | -| `max_outstanding` | `256` | `RadixClient` 未领取 job 的上限。 | - -### 2.3 不再接受的段 - -旧格式的 `cluster` / `data` / `index` 段出现时直接报错:这些键现在都是 `radix-server` 的命令行参数(第 7 节有对照表)。 - -示例: - -```yaml -# /etc/flexkv/radixshmem.yaml -server: - name: /flexkv - ready_timeout_s: 900 -client: - prefetch_timeout_ms: 5000 - prefetch_max_inflight: 128 -``` - ---- +| `READY_TIMEOUT_S` | 600 | 等 server 可达且 ready 的总时长,覆盖 server 晚起、SlotStore prefault、集群 rendezvous;server 的 `--bootstrap-timeout` 不要超过它。 | +| `PREFETCH_TIMEOUT_MS` | 5000 | 一次 prefetch 拉取的服务端超时,到期后 job 以本地命中的部分完成。 | +| `PREFETCH_MAX_INFLIGHT` | 128 | 每个 KVTaskEngine 在飞的 peer 拉取上限,达到后新的 prefetch 跳过 peer 查询。 | +| `MAX_OUTSTANDING` | 256 | `RadixClient` 未领取 job 的上限。 | ## 3. 几何与 slot 数 @@ -82,53 +47,46 @@ FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `sh | 字段 | 来源 | |---|---| -| `block_size` | `CacheConfig.tokens_per_block`(sglang 的 page size) | -| `full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 CPU block 字节数(每 PP 段层数 × 节点内 KV head 数 × head_size × kv_dim × dtype × tokens_per_block) | -| `swa_slot_bytes` / `swa_window_blocks` | `CacheConfig.swa` 开启时:一个 SWA page 的字节数(uint8)与窗口块数;未开启则没有 SWA 池 | -| `slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,保证 SlotStore stride 等于 block 字节数 | -| `register_chunk_tokens` | 不传(0):RHT 注册粒度由 server 的 `--register-chunk-tokens` 决定(radixshmem 默认 4096 token)。FlexKV 不再有自己的 4096 常量,attach 后采纳 server 发布的值,并按 radixshmem 的规则换算成 block 数(`tokens // tokens_per_block`,至少 1):`adopt_geometry` 返回的 `register_chunk_tokens` / `register_chunk_blocks`,`CacheEngineRadixShmem` 的同名属性。`RadixGeometry.register_chunk_tokens` 非 0 时原样交给 server(pin) | - -server 收到几何后的规划(radixshmem 的规则):`swa_slots = floor(swa_ratio × data_bytes / swa_stride)`, -`full_slots = (data_bytes − SWA 占用) / full_stride`。任一池算出 0 个 slot、模型有 SWA 而 `--swa-ratio` 为 0, -都在 configure 时拒绝,FlexKV 报 `cannot serve FlexKV's geometry`。 +| `block_size` | `CacheConfig.tokens_per_block` | +| `full_slot_bytes` | 按 `StorageEngine` 的 BLOCKFIRST 布局算出的一个 CPU block 字节数 | +| `swa_slot_bytes` / `swa_window_blocks` | `CacheConfig.swa` 开启时一个 SWA page 的字节数与窗口块数;未开启则没有 SWA 池 | +| `slot_align` | 不超过 4096 且整除每个池 slot 字节数的最大二次幂,使 SlotStore stride 等于 block 字节数 | +| `register_chunk_tokens` | `0`,即 server 的 `--register-chunk-tokens`(默认 4096 token);采纳后按 `tokens // tokens_per_block`(至少 1)换算成 block 数 | -**采纳**:attach 成功后 `adopt_geometry`(KVManager 和 KVTaskEngine 通过 `adopt_radix_server` 调用)把 `pools.full.num_slots` 写进 `CacheConfig.num_cpu_blocks`, -`pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,并带回 server 的 `register_chunk_tokens` 及其 block 数,日志形如 -`adopted radix-server /flexkv's geometry: FULL 8605 slots (cpu_cache_gb had given 1524), SWA 1024 slots (had 1024); RHT registration chunk 4096 tokens = 64 blocks`。 -之后 TE 的 StorageEngine、cache engine、指标都用采纳后的值。 +server 的规划:`swa_slots = floor(swa_ratio × data_bytes / swa_stride)`,`full_slots = (data_bytes − SWA 占用) / full_stride`。 +任一池为 0、模型有 SWA 而 `--swa-ratio` 为 0,configure 时拒绝,FlexKV 报 `cannot serve FlexKV's geometry`。 -**校验**:每个 attach 方(KVManager、KVTaskEngine、cache engine、TE)用 `check_geometry` 复核 server 发布的 -`block_size`、各池 `slot_bytes`、SlotStore stride、SWA 窗口与自己的布局一致,不一致报错退出,不会静默错位传输。 -slot 数不在校验范围内,它们是 server 的;`register_chunk_tokens` 只在 FlexKV pin 了值时比对,server 的值不是 -`tokens_per_block` 的整数倍时只告警(chunk 取整到整 block)。 +**采纳**(`adopt_radix_server`,KVManager 与 KVTaskEngine 各调一次,幂等):`pools.full.num_slots` 写进 +`CacheConfig.num_cpu_blocks`,`pools.swa.num_slots` 写进 `CacheConfig.swa.num_slots`,同时记录 `register_chunk_tokens`。 +日志形如 `adopted radix-server /flexkv's geometry: FULL 8605 slots ..., SWA 1024 slots ...; RHT registration chunk 4096 tokens = 64 blocks`。 -**同一 server 上的多个 client** 必须带相同的几何:相同模型、page size、SWA 配置。第二个不同的几何被 server 以 -`GeometryMismatch` 拒绝,FlexKV 报 `already serves another geometry`。单机下 TP 不同的同一模型通常几何相同 -(节点内 KV head 数与 TP 无关),以 `check_geometry` 为准。 +**校验**(`check_geometry`,TE 等 attach 方):server 发布的 `block_size`、各池 `slot_bytes`、SlotStore stride、SWA 窗口 +须与自身布局一致,否则报错退出。slot 数是 server 的,不在校验范围。 ---- +同一 server 上的多个 client 须带相同的几何(相同模型、page size、SWA 配置);不同的几何被 server 以 `GeometryMismatch` +拒绝,FlexKV 报 `already serves another geometry`。 -## 4. 启动方式 +## 4. 启动 ### 4.1 单机 ```bash -# 运维,每节点一次;nohup / systemd 皆可。DSv4 这类有 SWA 池的模型给 --swa-ratio。 +# 运维,每节点一次;nohup / systemd 皆可。有 SWA 池的模型(DSv4 等)给 --swa-ratio。 radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 # 推理引擎侧 export FLEXKV_ENABLE_RADIXSHMEM=1 export FLEXKV_CPU_LAYOUT=BLOCKFIRST -# 不设 FLEXKV_RADIXSHMEM_CONFIG_PATH 即 attach /flexkv +# 不设 FLEXKV_RADIXSHMEM_SERVER_NAME 即 attach /flexkv ``` -server 起来后打印 `Waiting for a client's geometry`;FlexKV 的第一个进程 attach 时把几何交过去,server 建区域后 ready, -所有进程的 `wait_ready` 返回。server 晚于引擎启动也可以:FlexKV 在 `ready_timeout_s` 内重试连接。 +server 起来后打印 `Waiting for a client's geometry`;FlexKV 的第一个进程 attach 时交出几何,server 建区域后 ready。 +server 可晚于引擎启动,FlexKV 在 `ready_timeout_s` 内重试连接。 ### 4.2 多机(一个集群) -每个节点各起一个 server,用相同的 `--cluster-id` 和 `--registry`;`--rpc-interface`(或 `--rpc-address`)给对端拨入的 IP, -`--node-name` 空时自动为 `node`: +每个节点各起一个 server,相同的 `--cluster-id` 和 `--registry`;`--rpc-interface`(或 `--rpc-address`)给对端拨入的 IP, +`--node-name` 空时为 `node`: ```bash radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 \ @@ -138,71 +96,41 @@ radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 \ --transfer-dev mlx5_1 --transfer-dev mlx5_2 --bootstrap-timeout 600 ``` -集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`、`register_chunk_tokens`)由第一个拿到几何的节点发布到 -etcd `radix//geometry/`,其余 `waiting` 的节点采纳;各节点的 slot 数可以不同(预算可以不同)。 -FlexKV 侧每个节点同一份 YAML 即可。`ready_timeout_s` 要不小于 `--bootstrap-timeout`。 +集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`、`register_chunk_tokens`)由第一个 +拿到几何的节点发布到 etcd `radix//geometry/`,其余节点采纳;各节点的 slot 数可以不同。FlexKV 侧每个节点 +同一个 `FLEXKV_RADIXSHMEM_SERVER_NAME`。 ### 4.3 同机多节点(测试) -两个 server 在一台机器上:不同的 `--name`(或相同 name 加 `--endpoint`、`--data-name` 区分)、不同的 `--node-name`、 -`--rpc-address 127.0.0.1`。两个 FlexKV 进程各用一份 YAML,`server.name` / `server.endpoint` 指向自己的 server, -并各给一个 `FLEXKV_SERVER_RECV_PORT`。 +两个 server 在一台机器上:不同的 `--name`、不同的 `--node-name`、`--rpc-address 127.0.0.1`。两个 FlexKV 进程各自的 +`FLEXKV_RADIXSHMEM_SERVER_NAME` 指向自己的 server,并各给一个 `FLEXKV_SERVER_RECV_PORT`。 ### 4.4 一节点多引擎共享一个 server -两个独立的推理引擎(各自的 FlexKV、各自的 GPU)attach 同一个 radix-server,互相命中对方存的 KV。这就是 FlexKV -原有的多实例模式: +两个独立的推理引擎(各自的 FlexKV、各自的 GPU)attach 同一个 radix-server,互相命中对方存的 KV: ```bash # 引擎 A # 引擎 B FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=1 ``` -同一份 YAML。`instance_num > 1` 自动进入 server-client 模式:`instance 0` 的 dp0 内嵌本节点唯一的 KVServer(或者用 -`FLEXKV_SERVER_LAUNCH_MODE=external` 单独启动它),KVServer 里的 KVTaskEngine attach radix-server,它的 TE 等到 -`instance_num × gpus_per_node` 张 GPU 都注册才 ready,所以两个引擎都要启动。两边模型 / page size / SWA 配置必须相同 -(第 3 节)。node-local DP(多机 DP attention)路径下 `local_dp_client_id` 不带 instance,多实例暂不支持。 +同一个 `FLEXKV_RADIXSHMEM_SERVER_NAME`。`instance_num > 1` 走 server-client 模式:`instance 0` 的 dp0 内嵌本节点的 KVServer(或用 +`FLEXKV_SERVER_LAUNCH_MODE=external` 单独启动),它的 TE 等 `instance_num × gpus_per_node` 张 GPU 都注册后 ready, +所以两个引擎都要启动。两边模型 / page size / SWA 配置必须相同(第 3 节)。node-local DP(多机 DP attention)下不支持多实例。 ---- - -## 5. 命名派生 +## 5. 命名 | 对象 | 名字 | |---|---| | index shm | `--name`;集群模式下 radixshmem 追加 `_`,attach 方只需 `--name` | | SlotStore shm | `_data`(`--data-name` 可改) | -| gRPC socket | `/dev/shm/.sock`(`--endpoint` 可改;YAML `server.endpoint` 跟着改) | +| gRPC socket | `/dev/shm/.sock`,FlexKV 按名字派生;server 端保持默认 `--endpoint` | | etcd 键空间 | `radix//...` | ---- - -## 6. 启动时校验 - -FlexKV 加载 YAML 时报错的情况:出现 `cluster` / `data` / `index` 段;未知段或未知键;`server.name` 不以 `/` 开头或含空白; -`server.ready_timeout_s <= 0`;`client.prefetch_max_inflight >= client.max_outstanding`;`client.prefetch_timeout_ms <= 0`。 +## 6. 报错含义 -`CacheConfig` 侧:`FLEXKV_CPU_LAYOUT != BLOCKFIRST`;打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`; -SWA 开启但 `window_blocks < 1`。 - -attach 时:`ready_timeout_s` 内连不上 server 报 `no radix-server named ... reachable ...(start it with radix-server --name ...)`; -server 一直在等几何或配置失败报 `not ready within ...`(带 server 的 `mode` 和 `last_error`);几何冲突见第 3 节。 - ---- - -## 7. 从旧版迁移 - -| 旧 YAML 键 | 现在 | -|---|---| -| `cluster.cluster_id` | `radix-server --cluster-id`(同时不再派生 shm 名;shm 名是 `--name`) | -| `cluster.expected_min_nodes` / `num_rht_shards` / `rht_shard_holders` / `rht_slots_per_bucket` | `--expected-min-nodes` / `--num-rht-shards` / `--rht-shard-holders` / `--rht-slots` | -| `cluster.registry` / `rpc_interface` / `rpc_port` / `settle_ms` / `bootstrap_timeout_sec` | `--registry` / `--rpc-interface` / `--rpc-port` / `--settle-ms` / `--bootstrap-timeout` | -| `cluster.index_dev` / `gid_idx` / `rht_transport` / `peer_index_transport` / `remote_op_transport` / `zmq_listen_port` | 同名 `--index-dev` 等 | -| `data.transfer_devices` / `transfer_protocol` / `transfer_ip` / `transfer_port` / `transfer_metadata` | `--transfer-dev`(可重复)/ `--transfer-protocol` / `--transfer-ip` / `--transfer-port` / `--transfer-metadata` | -| `data.prefault` / `max_inflight` / `max_pending_jobs` / `job_ttl_s` | `--no-prefault` / `--max-inflight` / `--max-pending-jobs` / `--job-ttl` | -| `index.data_pool_ratio` / `background_evict_ratio` / `max_nodes` | `--data-pool-ratio` / `--background-evict-ratio` / `--max-nodes` | -| `index.register_chunk_size`(block 数;FlexKV 曾固定按 4096 token 换算) | `--register-chunk-tokens`(token 数,默认 4096);FlexKV 采纳 server 的值,见第 3 节 | -| `server.endpoint` | 保留:server 的 `--endpoint` 与 YAML `server.endpoint` 各写一次 | -| `server.rpc_workers` / `hugepage_path` | `--rpc-workers` / `--hugepage-path` | -| (由 FlexKV 推导的 slot 数、`data_bytes`) | slot 数由 `--data-bytes` 和 `--swa-ratio` 决定,FlexKV 采纳 | -| `FLEXKV_RADIX_SERVER_LAUNCH_MODE` | 移除,只有外部 server | -| `FLEXKV_RADIX_NODE_NAME` / `FLEXKV_RADIX_RPC_ADDRESS` | `--node-name` / `--rpc-address` | +- `FLEXKV_RADIXSHMEM_SERVER_NAME` 须以 `/` 开头且无空白。 +- `CacheConfig`:`FLEXKV_CPU_LAYOUT != BLOCKFIRST`;打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`; + SWA 开启但 `window_blocks < 1`。 +- attach:600 s 内连不上 server 报 `no radix-server named ... reachable ...`(附启动命令);server 一直在等几何或 + 配置失败报 `not ready within ...`(附 server 的 `mode` 和 `last_error`);几何冲突见第 3 节。 diff --git a/examples/radixshmem_configs/radix_server_multi_node.sh b/examples/radixshmem/radix_server_multi_node.sh similarity index 100% rename from examples/radixshmem_configs/radix_server_multi_node.sh rename to examples/radixshmem/radix_server_multi_node.sh diff --git a/examples/radixshmem_configs/radix_server_single_node.sh b/examples/radixshmem/radix_server_single_node.sh similarity index 100% rename from examples/radixshmem_configs/radix_server_single_node.sh rename to examples/radixshmem/radix_server_single_node.sh diff --git a/examples/radixshmem_configs/radixshmem.yaml b/examples/radixshmem_configs/radixshmem.yaml deleted file mode 100644 index 6a8279514..000000000 --- a/examples/radixshmem_configs/radixshmem.yaml +++ /dev/null @@ -1,19 +0,0 @@ -# FlexKV side of radixshmem mode: which radix-server to attach to and the -# prefetch limits. Every key is at its default; the file is equivalent to not -# setting FLEXKV_RADIXSHMEM_CONFIG_PATH at all. -# -# export FLEXKV_ENABLE_RADIXSHMEM=1 -# export FLEXKV_CPU_LAYOUT=BLOCKFIRST -# export FLEXKV_RADIXSHMEM_CONFIG_PATH=$PWD/examples/radixshmem_configs/radixshmem.yaml -# -# The server itself is the operator's process (radix_server_single_node.sh / -# radix_server_multi_node.sh). FlexKV hands it the geometry and adopts the slot -# counts it plans from --data-bytes. Reference: docs/radixshmem/config_zh.md -server: - name: /flexkv # radix-server --name - endpoint: "" # "" = unix:///dev/shm/flexkv.sock - ready_timeout_s: 600 # reachable AND ready (prefault, cluster rendezvous) -client: - prefetch_timeout_ms: 5000 - prefetch_max_inflight: 128 - max_outstanding: 256 diff --git a/flexkv/cache/radix_shmem_planner.py b/flexkv/cache/radix_shmem_planner.py index 5f95852b5..8c65eb83b 100644 --- a/flexkv/cache/radix_shmem_planner.py +++ b/flexkv/cache/radix_shmem_planner.py @@ -51,7 +51,8 @@ from flexkv.common.block import SequenceMeta from flexkv.common.config import CacheConfig, ModelConfig from flexkv.common.debug import flexkv_logger -from flexkv.common.radixshmem_config import RadixShmemConfig, get_radixshmem_config +from flexkv.server.shm_radix_bootstrap import (PREFETCH_MAX_INFLIGHT, PREFETCH_TIMEOUT_MS, + expected_geometry, radix_server_name) from flexkv.common.transfer import ( DeviceType, TransferOp, @@ -160,9 +161,8 @@ def __init__(self, redis_meta=None, event_collector: Optional[KVEventCollector] = None): _check_cache_config(cache_config) - self._radix_config: RadixShmemConfig = get_radixshmem_config() # GetJobs this engine started and has not yet seen finish; pruned on - # every prefetch and used for back-pressure (client.prefetch_max_inflight). + # every prefetch and used for back-pressure (PREFETCH_MAX_INFLIGHT). self._prefetch_jobs: List[Any] = [] super().__init__(cache_config, model_config, redis_meta, event_collector) @@ -179,11 +179,8 @@ def _build_cpu_cache_engine(self, counts into `cache_config`) and waits for the server to be ready. Peer reuse follows the server: on whenever it is part of a cluster. """ - from flexkv.server.shm_radix_bootstrap import expected_geometry - - rcfg = self._radix_config return CacheEngineRadixShmem( - rcfg.server_name, + radix_server_name(), geometry=expected_geometry(self.model_config, cache_config), tokens_per_block=cache_config.tokens_per_block, num_total_blocks=cache_config.num_cpu_blocks, @@ -433,12 +430,11 @@ def _plan_prefetch(self, engine = self.cpu_cache_engine if not engine.peer_enabled: return RadixGetPlan.empty() - client_settings = self._radix_config.client inflight = self._prefetch_inflight() - if inflight >= client_settings.prefetch_max_inflight: + if inflight >= PREFETCH_MAX_INFLIGHT: flexkv_logger.debug( f"radixshmem prefetch {request_id}: {inflight} peer pulls in flight " - f"(limit {client_settings.prefetch_max_inflight}); skipping the peer walk") + f"(limit {PREFETCH_MAX_INFLIGHT}); skipping the peer walk") return RadixGetPlan.empty() swa_active = swa_aware and self.swa_op_constructor.enabled mask = (COMPONENT_MASK_FULL | COMPONENT_MASK_SWA) if swa_active else COMPONENT_MASK_FULL @@ -446,7 +442,7 @@ def _plan_prefetch(self, sequence_meta, component_mask=mask, query_end=block_mask_end, - timeout_ms=client_settings.prefetch_timeout_ms, + timeout_ms=PREFETCH_TIMEOUT_MS, ) plan = RadixGetPlan.empty() if job is None: diff --git a/flexkv/common/config.py b/flexkv/common/config.py index ae9183e12..e086416f5 100644 --- a/flexkv/common/config.py +++ b/flexkv/common/config.py @@ -839,14 +839,13 @@ def __str__(self) -> str: server_launch_mode=os.getenv('FLEXKV_SERVER_LAUNCH_MODE', 'embedded').lower(), server_recv_port=os.getenv('FLEXKV_SERVER_RECV_PORT', 'ipc:///tmp/flexkv_server'), - # radixshmem mode: the CPU tier is radixshmem's index + SlotStore, one - # radix-server per node (a process the operator starts: `radix-server - # --name /flexkv --data-bytes ...`), one shared TE, a KVTaskEngine per DP - # process (no KVServer). FlexKV only attaches to the server; which one and - # the prefetch limits are the YAML at FLEXKV_RADIXSHMEM_CONFIG_PATH - # (flexkv.common.radixshmem_config; reference docs/radixshmem/config_zh.md). + # radixshmem mode: the CPU tier is a radix-server (index + SlotStore), one + # per node, a process the operator starts (`radix-server --name /flexkv + # --data-bytes ...`). FlexKV attaches to it by name; everything else about + # the attach is a fixed default (flexkv.server.shm_radix_bootstrap; + # reference docs/radixshmem/config_zh.md). enable_radixshmem=bool(int(os.getenv('FLEXKV_ENABLE_RADIXSHMEM', 0))), - radixshmem_config_path=os.getenv('FLEXKV_RADIXSHMEM_CONFIG_PATH', '') or None, + radixshmem_server_name=os.getenv('FLEXKV_RADIXSHMEM_SERVER_NAME', '') or '/flexkv', index_accel=bool(int(os.getenv('FLEXKV_INDEX_ACCEL', 1))), cpu_layout_type=KVCacheLayoutType(os.getenv('FLEXKV_CPU_LAYOUT', 'BLOCKFIRST').upper()), diff --git a/flexkv/common/radixshmem_config.py b/flexkv/common/radixshmem_config.py deleted file mode 100644 index 10a3efe5e..000000000 --- a/flexkv/common/radixshmem_config.py +++ /dev/null @@ -1,240 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# cython: boundscheck=True, wraparound=True -"""The radixshmem-mode configuration file (``FLEXKV_RADIXSHMEM_CONFIG_PATH``). - -In radixshmem mode the CPU tier is a ``radix-server`` process the operator -starts on every node (``radix-server --name /flexkv --data-bytes 64G ...``): -the index shm, the SlotStore, the transfer engine and the cluster membership -all belong to that process and are set on its command line. FlexKV never -creates a server; it attaches a ``shmradix.RadixClient``. So this file holds -the two things FlexKV has to know, and nothing else: - - server - Which radix-server to attach to: its ``--name`` (which also derives the - default gRPC socket ``unix:///dev/shm/.sock``), an ``endpoint`` - override, and how long a FlexKV process waits for the server to exist - and become ready. - client - FlexKV's RadixClient / prefetch settings. - -The slot geometry (tokens per block, bytes of one CPU block and one SWA page, -the SWA window) is derived from ``ModelConfig`` / ``CacheConfig`` and handed to -the server by FlexKV's clients (``flexkv.server.shm_radix_bootstrap``); the -slot COUNTS come back from the server, which plans them from its byte budget. -None of that is in this file. A file that still carries the former ``cluster`` -/ ``data`` / ``index`` sections is rejected with a pointer to the -``radix-server`` flags they moved to. - -No YAML at all is a valid configuration: it attaches to ``radix-server --name -/flexkv`` on the local socket. - -Reference: ``docs/radixshmem/config_zh.md``. -""" -from __future__ import annotations - -import dataclasses -import threading -from typing import Any, Dict, Optional, Tuple - -import yaml - -from flexkv.common.config import GLOBAL_CONFIG_FROM_ENV - -SECTIONS = ("server", "client") -# Sections of the previous file format. Their keys are radix-server flags now. -RETIRED_SECTIONS = ("cluster", "data", "index") - -DEFAULT_SERVER_NAME = "/flexkv" - - -class RadixShmemConfigError(ValueError): - """The file is not a valid radixshmem-mode configuration.""" - - -def default_endpoint(server_name: str) -> str: - """radixshmem's default gRPC socket for ``radix-server --name ``: - ``unix:///dev/shm/ '_'>.sock``.""" - return f"unix:///dev/shm/{server_name.lstrip('/').replace('/', '_')}.sock" - - -@dataclasses.dataclass(frozen=True) -class RadixServerSettings: - """Which radix-server this node's FlexKV attaches to.""" - # ``radix-server --name``: the index shm name. Also the default socket - # (``unix:///dev/shm/.sock``, ``default_endpoint``). - name: str = DEFAULT_SERVER_NAME - # gRPC endpoint; "" = the default socket derived from ``name``. - endpoint: str = "" - # How long a FlexKV process waits for the server to be reachable AND ready. - # Covers the operator starting it late, the SlotStore prefault and, on a - # cluster, the rendezvous (the server's --bootstrap-timeout). - ready_timeout_s: float = 600.0 - - -@dataclasses.dataclass(frozen=True) -class RadixClientSettings: - """FlexKV-side settings of the RadixClient and the prefetch path.""" - # Server-side deadline of one prefetch pull; the job completes with the - # local hit when it expires. - prefetch_timeout_ms: int = 5000 - # Peer pulls in flight per CE process before new prefetches skip the peer - # walk; kept under max_outstanding so pull_async never blocks. - prefetch_max_inflight: int = 128 - # Uncollected jobs one RadixClient may hold. - max_outstanding: int = 256 - - -@dataclasses.dataclass(frozen=True) -class RadixShmemConfig: - path: Optional[str] - server: RadixServerSettings = RadixServerSettings() - client: RadixClientSettings = RadixClientSettings() - - # ------------------------------------------------------------ server - @property - def server_name(self) -> str: - return self.server.name - - @property - def endpoint(self) -> str: - """gRPC endpoint; "" = radixshmem's ``unix:///dev/shm/.sock``.""" - return self.server.endpoint - - @property - def ready_timeout_s(self) -> float: - return float(self.server.ready_timeout_s) - - @property - def default_endpoint(self) -> str: - """The gRPC endpoint radixshmem derives from the server name when no - ``endpoint`` is given: ``unix:///dev/shm/.sock``.""" - return default_endpoint(self.server.name) - - # ------------------------------------------------------------- tests - def replace_server(self, **changes: Any) -> "RadixShmemConfig": - return dataclasses.replace(self, server=dataclasses.replace(self.server, **changes)) - - def replace_client(self, **changes: Any) -> "RadixShmemConfig": - return dataclasses.replace(self, client=dataclasses.replace(self.client, **changes)) - - def describe(self) -> str: - where = self.path or "(defaults)" - return (f"{where}: radix-server {self.server_name} " - f"(endpoint={self.endpoint or self.default_endpoint}, " - f"ready_timeout_s={self.ready_timeout_s:.0f})") - - -# ------------------------------------------------------------------ loading - -def _read_yaml(path: str) -> Dict[str, Any]: - with open(path) as f: - loaded = yaml.safe_load(f) - if loaded is None: - return {} - if not isinstance(loaded, dict): - raise RadixShmemConfigError(f"{path}: top level must be a mapping of sections") - return loaded - - -def _section(raw: Dict[str, Any], name: str, path: str) -> Dict[str, Any]: - sec = raw.get(name) - if sec is None: - return {} - if not isinstance(sec, dict): - raise RadixShmemConfigError(f"{path}: section '{name}' must be a mapping") - return dict(sec) - - -def _typed_section(name: str, given: Dict[str, Any], dc, path: str): - """``dc(**given)`` after checking the keys and coercing the value types.""" - fields = {f.name: f for f in dataclasses.fields(dc)} - unknown = set(given) - set(fields) - if unknown: - raise RadixShmemConfigError( - f"{path}: unknown key(s) in '{name}': {sorted(unknown)}; expected {sorted(fields)}") - values: Dict[str, Any] = {} - for key, value in given.items(): - typ = fields[key].type - try: - if typ in ("int", int): - values[key] = int(value) - elif typ in ("float", float): - values[key] = float(value) - else: - values[key] = "" if value is None else str(value) - except (TypeError, ValueError) as exc: - raise RadixShmemConfigError(f"{path}: '{name}.{key}' has an invalid value {value!r}") from exc - return dc(**values) - - -def _validate(cfg: RadixShmemConfig, path: str) -> None: - name = cfg.server_name - if not name or not name.startswith("/") or len(name) < 2 or any(c.isspace() for c in name): - raise RadixShmemConfigError( - f"{path}: server.name={name!r} must be a shm name that starts with '/' " - f"(the radix-server's --name, e.g. '/flexkv')") - if cfg.ready_timeout_s <= 0: - raise RadixShmemConfigError(f"{path}: server.ready_timeout_s must be > 0") - if cfg.client.prefetch_max_inflight >= cfg.client.max_outstanding: - raise RadixShmemConfigError( - f"{path}: client.prefetch_max_inflight={cfg.client.prefetch_max_inflight} must be " - f"below client.max_outstanding={cfg.client.max_outstanding}, or pull_async blocks") - if cfg.client.prefetch_timeout_ms <= 0: - raise RadixShmemConfigError(f"{path}: client.prefetch_timeout_ms must be > 0") - - -def load_radixshmem_config(path: Optional[str] = None) -> RadixShmemConfig: - """Parse ``path`` (None or "" = all defaults: ``radix-server --name /flexkv`` - on the local socket). Raises :class:`RadixShmemConfigError` on an invalid - file.""" - label = path or "(defaults)" - raw = _read_yaml(path) if path else {} - retired = [s for s in RETIRED_SECTIONS if s in raw] - if retired: - raise RadixShmemConfigError( - f"{label}: section(s) {retired} are not FlexKV's any more: the radix-server owns " - f"its cluster, data plane and index settings and takes them on its command line " - f"(radix-server --data-bytes / --swa-ratio / --expected-min-nodes / --registry / " - f"--transfer-dev ...). FlexKV only attaches to it; keep 'server' and 'client' here. " - f"See docs/radixshmem/config_zh.md") - unknown = set(raw) - set(SECTIONS) - if unknown: - raise RadixShmemConfigError( - f"{label}: unknown section(s) {sorted(unknown)}; expected {list(SECTIONS)}") - cfg = RadixShmemConfig( - path=path or None, - server=_typed_section("server", _section(raw, "server", label), RadixServerSettings, label), - client=_typed_section("client", _section(raw, "client", label), RadixClientSettings, label), - ) - _validate(cfg, label) - return cfg - - -# -------------------------------------------------------------- singleton - -_lock = threading.Lock() -_cached: Optional[Tuple[Tuple[str], RadixShmemConfig]] = None - - -def _env_key() -> Tuple[str]: - return (str(GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path or ""),) - - -def get_radixshmem_config() -> RadixShmemConfig: - """The process's configuration: loaded from ``GLOBAL_CONFIG_FROM_ENV`` - (``FLEXKV_RADIXSHMEM_CONFIG_PATH``) on first use and whenever that value - changes.""" - global _cached - key = _env_key() - with _lock: - if _cached is None or _cached[0] != key: - _cached = (key, load_radixshmem_config(key[0] or None)) - return _cached[1] - - -def set_radixshmem_config(cfg: Optional[RadixShmemConfig]) -> None: - """Install ``cfg`` as the process's configuration (tests); None reverts to - loading from the environment.""" - global _cached - with _lock: - _cached = None if cfg is None else (_env_key(), cfg) diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 642676668..58b5d62e0 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -41,8 +41,6 @@ from flexkv.common.config import (GLOBAL_CONFIG_FROM_ENV, CacheConfig, LayerGroupSpec, ModelConfig, SWAPoolConfig) from flexkv.common.debug import flexkv_logger -from flexkv.common.radixshmem_config import (RadixShmemConfig, default_endpoint, - get_radixshmem_config) from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType try: @@ -54,6 +52,39 @@ # Pool bases are page aligned regardless; a larger per-slot alignment only pads. _MAX_SLOT_ALIGN = 4096 +# FlexKV's side of the attach. The server to attach to is the one setting +# (FLEXKV_RADIXSHMEM_SERVER_NAME, radix_server_name()); the rest are fixed. +DEFAULT_SERVER_NAME = "/flexkv" +# How long a FlexKV process waits for the server to be reachable AND ready: +# covers the operator starting it late, the SlotStore prefault and, on a +# cluster, the rendezvous (keep the server's --bootstrap-timeout below this). +READY_TIMEOUT_S = 600.0 +# Server-side deadline of one peer pull; the job completes with the local hit +# when it expires. +PREFETCH_TIMEOUT_MS = 5000 +# Peer pulls in flight per KVTaskEngine before new prefetches skip the peer +# walk; kept below MAX_OUTSTANDING so pull_async never blocks. +PREFETCH_MAX_INFLIGHT = 128 +# Uncollected jobs one RadixClient may hold (radixshmem's own default). +MAX_OUTSTANDING = 256 + + +def radix_server_name() -> str: + """The ``--name`` of this node's radix-server: ``FLEXKV_RADIXSHMEM_SERVER_NAME`` + (default ``/flexkv``). A shm name: starts with '/', no whitespace.""" + name = str(GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name or DEFAULT_SERVER_NAME) + if not name.startswith("/") or len(name) < 2 or any(c.isspace() for c in name): + raise ValueError( + f"FLEXKV_RADIXSHMEM_SERVER_NAME={name!r} must be a shm name that starts with '/' " + f"(the radix-server's --name, e.g. '/flexkv')") + return name + + +def default_endpoint(server_name: str) -> str: + """The gRPC socket radixshmem derives from ``radix-server --name ``: + ``unix:///dev/shm/ '_'>.sock``.""" + return f"unix:///dev/shm/{server_name.lstrip('/').replace('/', '_')}.sock" + def _ensure_shmradix() -> None: if shmradix is None: @@ -238,13 +269,11 @@ def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> R def attach_radix_client(name: Optional[str] = None, *, geometry: Any = None, - rcfg: Optional[RadixShmemConfig] = None, - endpoint: Optional[str] = None, timeout_s: Optional[float] = None, max_outstanding: Optional[int] = None, label: str = "radixshmem") -> "shmradix.RadixClient": """A ready ``shmradix.RadixClient`` on the radix-server ``name`` (default: - the configuration's ``server.name``). + :func:`radix_server_name`), at the socket radixshmem derives from the name. With ``geometry`` (a :class:`RadixGeometry` or a ``shmradix.Geometry``) the client hands the server FlexKV's slot shape on the way; the server plans @@ -257,27 +286,22 @@ def attach_radix_client(name: Optional[str] = None, Retries while the server is not reachable yet (the operator may start it late), then blocks in ``wait_ready`` -- the rendezvous of a cluster and the SlotStore prefault happen there -- for ``timeout_s`` in total - (default: the configuration's ``server.ready_timeout_s``). + (default ``READY_TIMEOUT_S``). """ _ensure_shmradix() - if rcfg is None: - rcfg = get_radixshmem_config() - name = name or rcfg.server_name - if endpoint is None: - endpoint = rcfg.endpoint or None + name = name or radix_server_name() if timeout_s is None: - timeout_s = rcfg.ready_timeout_s + timeout_s = READY_TIMEOUT_S if max_outstanding is None: - max_outstanding = rcfg.client.max_outstanding + max_outstanding = MAX_OUTSTANDING spec = geometry.to_shmradix() if isinstance(geometry, RadixGeometry) else geometry - where = endpoint or default_endpoint(name) + where = default_endpoint(name) deadline = time.monotonic() + float(timeout_s) last: Optional[BaseException] = None while True: try: - client = shmradix.RadixClient(name, spec, endpoint=endpoint, - max_outstanding=max_outstanding) + client = shmradix.RadixClient(name, spec, max_outstanding=max_outstanding) break except shmradix.GeometryMismatch as e: raise ValueError( @@ -440,7 +464,7 @@ def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, *, - rcfg: Optional[RadixShmemConfig] = None, + name: Optional[str] = None, label: str = "radixshmem") -> Dict[str, int]: """Attach to this node's radix-server with FlexKV's geometry, take over the slot counts it planned (:func:`adopt_geometry`) and its cluster rank @@ -451,7 +475,7 @@ def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, the TE. Idempotent: the server accepts the same geometry any number of times.""" geometry = expected_geometry(model_config, cache_config) - client = attach_radix_client(rcfg=rcfg, geometry=geometry, label=label) + client = attach_radix_client(name, geometry=geometry, label=label) try: counts = adopt_geometry(cache_config, client, label=label) cache_config.distributed_node_id = radix_cluster_rank(client) @@ -464,7 +488,7 @@ def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, client.close() -def radix_server_is_distributed(rcfg: Optional[RadixShmemConfig] = None, +def radix_server_is_distributed(name: Optional[str] = None, *, timeout_s: Optional[float] = None, label: str = "radixshmem") -> bool: @@ -473,7 +497,7 @@ def radix_server_is_distributed(rcfg: Optional[RadixShmemConfig] = None, is ready (a geometry-less attach that only waits), so every process that asks gets the same answer -- the framework adapters gate the prefetch path on it in every TP rank, and the ranks must agree.""" - client = attach_radix_client(rcfg=rcfg, timeout_s=timeout_s, label=label) + client = attach_radix_client(name, timeout_s=timeout_s, label=label) try: return int(client.info.world_size) > 1 finally: diff --git a/tests/radixshmem/radix_e2e_common.py b/tests/radixshmem/radix_e2e_common.py index deb7f68b6..d3c8da81a 100644 --- a/tests/radixshmem/radix_e2e_common.py +++ b/tests/radixshmem/radix_e2e_common.py @@ -92,16 +92,6 @@ def stop_private_etcd(proc, workdir) -> None: shutil.rmtree(workdir, ignore_errors=True) -def write_radix_config(workdir: str, config: dict, name: str = "radixshmem.yaml") -> str: - """Write the run's radixshmem YAML (``FLEXKV_RADIXSHMEM_CONFIG_PATH``) and - return its path.""" - import yaml - path = os.path.join(workdir, name) - with open(path, "w") as f: - yaml.safe_dump(config, f) - return path - - def start_radix_server(name: str, data_bytes: int, *, extra_args=(), endpoint: Optional[str] = None, log_path: Optional[str] = None, timeout: float = 60.0) -> subprocess.Popen: """Start the operator's ``radix-server`` (``python -m shmradix.cli``) and wait diff --git a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py index e35ac7bcb..454b8b76f 100644 --- a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -8,8 +8,7 @@ SlotStore = the node's CPU pool + RDMA transfer engine), joined into one cluster by their command lines (--expected-min-nodes 2, --registry, --node-name, RDMA devices); each FlexKV node attaches to its own server - through a per-node YAML (FLEXKV_RADIXSHMEM_CONFIG_PATH: server.name / - server.endpoint) and brings the geometry; the two servers rendezvous in + (FLEXKV_RADIXSHMEM_SERVER_NAME) and brings the geometry; the two servers rendezvous in one etcd namespace, get dense cluster ranks and an RHT to route by; * node 0 PUTs a window of GPU blocks holding a per-block pattern; * node 1 calls ``KVManager.prefetch_async`` for the same tokens: the index walk @@ -60,7 +59,6 @@ sweep_radix_files, wait_kv_manager_ready, write_pattern, - write_radix_config, ) WORLD_SIZE = 2 @@ -106,7 +104,7 @@ def _prefetch_until(kvm, token_ids, want_pulled_blocks: int, timeout: float = 60 return pulled, rounds -def _node_proc(rank, gpu_id, cluster_id, config_path, +def _node_proc(rank, gpu_id, cluster_id, server_name, reader_ready, written, read_done, result_q): """One FlexKV node: rank 0 writes the windows, rank 1 prefetches and reads.""" # Before any CUDA context exists: each node drives a different device while @@ -117,7 +115,7 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, recv_port = f"ipc:///tmp/flexkv_{cluster_id}_{node_name}" os.environ.update({ "FLEXKV_ENABLE_RADIXSHMEM": "1", - "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, # this node's server + "FLEXKV_RADIXSHMEM_SERVER_NAME": server_name, # this node's server "FLEXKV_ENABLE_MPS": "0", "FLEXKV_SERVER_RECV_PORT": recv_port, }) @@ -128,7 +126,7 @@ def _node_proc(rank, gpu_id, cluster_id, config_path, # Built from env at import time; set the fields that matter explicitly in # case a parent import happened earlier in this process. GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True - GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path = config_path + GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name = server_name GLOBAL_CONFIG_FROM_ENV.enable_mps = False GLOBAL_CONFIG_FROM_ENV.server_recv_port = recv_port @@ -240,21 +238,18 @@ def _run(registry: str, rdma_dev: str) -> dict: CacheConfig(tokens_per_block=TOKENS_PER_BLOCK, enable_cpu=True, enable_ssd=False, num_cpu_blocks=NUM_CPU_BLOCKS)) sweep_radix_files(cluster_id) - servers, config_paths = [], [] + servers, names = [], [] for rank in range(WORLD_SIZE): name = f"/{cluster_id}_{_node_name(rank)}" - endpoint = f"unix:///dev/shm/{cluster_id}_{_node_name(rank)}.sock" + names.append(name) servers.append(start_radix_server( - name, NUM_CPU_BLOCKS * block_bytes, endpoint=endpoint, + name, NUM_CPU_BLOCKS * block_bytes, extra_args=["--expected-min-nodes", str(WORLD_SIZE), "--registry", registry, "--cluster-id", cluster_id, "--node-name", _node_name(rank), "--rpc-address", "127.0.0.1", "--index-dev", rdma_dev, "--transfer-dev", rdma_dev, "--rht-slots", "4", "--bootstrap-timeout", "120"], log_path=os.path.join(workdir, f"radix-server-{_node_name(rank)}.log"))) - config_paths.append(write_radix_config( - workdir, {"server": {"name": name, "endpoint": endpoint, "ready_timeout_s": 300}}, - name=f"radixshmem_{_node_name(rank)}.yaml")) ctx = mp.get_context("spawn") reader_ready, written, read_done = ctx.Event(), ctx.Event(), ctx.Event() result_q = ctx.Queue() @@ -264,7 +259,7 @@ def _run(registry: str, rdma_dev: str) -> dict: for rank in range(WORLD_SIZE): proc = ctx.Process( target=_node_proc, - args=(rank, rank, cluster_id, config_paths[rank], + args=(rank, rank, cluster_id, names[rank], reader_ready, written, read_done, result_q), daemon=False, ) diff --git a/tests/radixshmem/test_e2e_radix_shmem.py b/tests/radixshmem/test_e2e_radix_shmem.py index 7418857e3..2fb237508 100644 --- a/tests/radixshmem/test_e2e_radix_shmem.py +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -9,8 +9,8 @@ node's CPU KV pool) with nothing but a name and a byte budget; every KVManager hands it FlexKV's geometry and adopts the slot counts it plans, and so does the process that builds the KVTaskEngine. Which server to - attach to is the ``server.name`` of a small YAML written per run - (FLEXKV_RADIXSHMEM_CONFIG_PATH), which is how a deployment names it too. + attach to is FLEXKV_RADIXSHMEM_SERVER_NAME, set per run, which is how a + deployment names it too. * Phase 1: every DP PUTs its own requests concurrently. * Phase 2 (dp_size > 1): dp0 PUTs a prefix that dp1 then finds with ``get_match`` -- the shared index is what the radixshmem path exists for. @@ -52,7 +52,6 @@ sweep_radix_files, wait_kv_manager_ready, write_pattern, - write_radix_config, ) NUM_GPU_BLOCKS = 256 @@ -65,7 +64,7 @@ ROUNDTRIP_START_BLOCK = 192 -def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, +def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, server_name: str, barrier, result_q) -> None: """Full lifecycle of one DP scheduler process.""" # Before the first flexkv import: GLOBAL_CONFIG_FROM_ENV is read at import. @@ -74,7 +73,7 @@ def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, recv_port = f"ipc:///tmp/flexkv_{server_id}" os.environ.update({ "FLEXKV_ENABLE_RADIXSHMEM": "1", - "FLEXKV_RADIXSHMEM_CONFIG_PATH": config_path, + "FLEXKV_RADIXSHMEM_SERVER_NAME": server_name, "FLEXKV_ENABLE_MPS": "0", "FLEXKV_SERVER_RECV_PORT": recv_port, }) @@ -84,7 +83,7 @@ def _dp_proc(dp_client_id: int, dp_size: int, server_id: str, config_path: str, from flexkv.kvmanager import KVManager GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True - GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path = config_path + GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name = server_name GLOBAL_CONFIG_FROM_ENV.enable_mps = False GLOBAL_CONFIG_FROM_ENV.server_recv_port = recv_port @@ -188,7 +187,6 @@ def _run(dp_size: int) -> dict: server_id = f"e2e{dp_size}dp_{os.getpid()}" workdir = tempfile.mkdtemp(prefix="flexkv_radix_e2e_") name = f"/{server_id}" - config_path = write_radix_config(workdir, {"server": {"name": name, "ready_timeout_s": 300}}) # The operator's server: a byte budget that holds NUM_CPU_BLOCKS blocks of # this test's model (the DP processes bring the geometry and adopt the count). from flexkv.common.config import CacheConfig, ModelConfig @@ -206,7 +204,7 @@ def _run(dp_size: int) -> dict: result_q = ctx.Queue() procs = [ ctx.Process(target=_dp_proc, - args=(dp, dp_size, server_id, config_path, barrier, result_q), + args=(dp, dp_size, server_id, name, barrier, result_q), daemon=False) for dp in range(dp_size) ] diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index b40c8a399..556c76020 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -102,8 +102,6 @@ def _load_module_direct(name: str, path: str): # Pure-Python (no c_ext): the bootstrap (server config, geometry, attach) and # the transfer enums. from flexkv.common.config import GLOBAL_CONFIG_FROM_ENV # noqa: E402 -from flexkv.common.radixshmem_config import ( # noqa: E402 - RadixShmemConfigError, load_radixshmem_config, set_radixshmem_config) from flexkv.common.transfer import TransferType # noqa: E402 from flexkv.server import shm_radix_bootstrap as bootstrap # noqa: E402 @@ -205,21 +203,16 @@ def close(self) -> None: self._stack.pop()() -def _radix_config(**server): - """The all-defaults radixshmem configuration with ``server`` keys changed; - a short ready timeout so a broken test fails instead of waiting.""" - return load_radixshmem_config(None).replace_server(**{"ready_timeout_s": 60.0, **server}) - - @pytest.fixture def env(): - set_radixshmem_config(_radix_config()) + saved = bootstrap.READY_TIMEOUT_S + bootstrap.READY_TIMEOUT_S = 60.0 # a broken test fails instead of waiting ten minutes e = _Env() try: yield e finally: e.close() - set_radixshmem_config(None) + bootstrap.READY_TIMEOUT_S = saved # ============================================================================= @@ -689,7 +682,7 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): client, dataclasses.replace(geo, register_chunk_tokens=4096), "test") assert bootstrap.radix_cluster_rank(client) == 0 and client.info.world_size == 1 # the connectors' prefetch gate asks the server the same question - assert bootstrap.radix_server_is_distributed(_radix_config(name=name), timeout_s=30) is False + assert bootstrap.radix_server_is_distributed(name, timeout_s=30) is False # another expectation against the same regions fails closed cache_config.tokens_per_block = 32 with pytest.raises(ValueError, match="tokens_per_block"): @@ -732,84 +725,30 @@ def _start_later(): bootstrap.attach_radix_client(f"/nobody{os.getpid()}", geometry=geo, timeout_s=2) -# ----------------------------------------------------------------------------- -# Part 1c — the radixshmem-mode YAML (flexkv.common.radixshmem_config): which -# radix-server to attach to and FlexKV's client settings; the former -# server-side sections are refused (docs/radixshmem/config_zh.md). - - -def _write_yaml(tmp_path, text: str) -> str: - path = tmp_path / "radixshmem.yaml" - path.write_text(text) - return str(path) - - -def test_radix_config_defaults(): - cfg = load_radixshmem_config(None) - assert cfg.path is None - assert cfg.server_name == "/flexkv" and cfg.endpoint == "" and cfg.ready_timeout_s == 600.0 - assert cfg.default_endpoint == "unix:///dev/shm/flexkv.sock" - assert cfg.client.prefetch_timeout_ms == 5000 and cfg.client.prefetch_max_inflight == 128 - assert cfg.client.max_outstanding == 256 - assert "radix-server /flexkv" in cfg.describe() - - -def test_radix_config_file(tmp_path): - path = _write_yaml(tmp_path, """ -server: - name: /prod/kv - endpoint: 10.0.0.2:7000 - ready_timeout_s: 900 -client: - prefetch_timeout_ms: 1000 - max_outstanding: 512 - prefetch_max_inflight: 300 -""") - cfg = load_radixshmem_config(path) - assert cfg.path == path and cfg.server_name == "/prod/kv" - assert cfg.default_endpoint == "unix:///dev/shm/prod_kv.sock" - assert cfg.endpoint == "10.0.0.2:7000" and cfg.ready_timeout_s == 900.0 - assert cfg.client.prefetch_timeout_ms == 1000 and cfg.client.max_outstanding == 512 - assert cfg.client.prefetch_max_inflight == 300 - assert cfg.replace_server(name="/x").server_name == "/x" - assert cfg.replace_client(max_outstanding=1000).client.max_outstanding == 1000 - - -@pytest.mark.parametrize("text, match", [ - ("cluster:\n cluster_id: prod\n", "radix-server"), # former server-side sections - ("data:\n prefault: false\n", "radix-server"), - ("index:\n data_pool_ratio: 8\n", "radix-server"), - ("peers: {}\n", "unknown section"), - ("- a\n", "must be a mapping"), - ("server: 5\n", "must be a mapping"), - ("server:\n rpc_workers: 3\n", "unknown key"), - ("server:\n name: kv\n", "starts with"), - ("server:\n name: /a b\n", "starts with"), - ("server:\n ready_timeout_s: 0\n", "ready_timeout_s"), - ("server:\n ready_timeout_s: soon\n", "invalid value"), - ("client:\n prefetch_max_inflight: 256\n", "max_outstanding"), - ("client:\n prefetch_timeout_ms: 0\n", "prefetch_timeout_ms"), - ("client:\n timeout: 5\n", "unknown key"), -]) -def test_radix_config_rejects(tmp_path, text, match): - with pytest.raises(RadixShmemConfigError, match=match): - load_radixshmem_config(_write_yaml(tmp_path, text)) - - -def test_radix_config_env_singleton_reloads_on_change(tmp_path, monkeypatch): - """`get_radixshmem_config` follows GLOBAL_CONFIG_FROM_ENV.radixshmem_config_path; - a test-installed config wins until reverted.""" - from flexkv.common.radixshmem_config import get_radixshmem_config - set_radixshmem_config(None) - monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", None) - assert get_radixshmem_config().server_name == "/flexkv" - path = _write_yaml(tmp_path, "server:\n name: /other\n") - monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_config_path", path) - assert get_radixshmem_config().server_name == "/other" - set_radixshmem_config(_radix_config(name="/pinned")) - assert get_radixshmem_config().server_name == "/pinned" - set_radixshmem_config(None) - assert get_radixshmem_config().default_endpoint == "unix:///dev/shm/other.sock" +# ============================================================================= +# Part 1c — FLEXKV_RADIXSHMEM_SERVER_NAME, the one radixshmem-mode setting +# (shm_radix_bootstrap.radix_server_name); everything else about the attach +# is a fixed default. +# ============================================================================= + + +def test_radix_server_name_follows_the_env(monkeypatch): + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_server_name", "/flexkv") + assert bootstrap.radix_server_name() == "/flexkv" + assert bootstrap.default_endpoint("/flexkv") == "unix:///dev/shm/flexkv.sock" + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_server_name", "/prod/kv") + assert bootstrap.radix_server_name() == "/prod/kv" + assert bootstrap.default_endpoint("/prod/kv") == "unix:///dev/shm/prod_kv.sock" + # the fixed defaults stay consistent with each other + assert bootstrap.READY_TIMEOUT_S > 0 and bootstrap.PREFETCH_TIMEOUT_MS > 0 + assert 0 < bootstrap.PREFETCH_MAX_INFLIGHT < bootstrap.MAX_OUTSTANDING + + +@pytest.mark.parametrize("name", ["kv", "/a b", "/"]) +def test_radix_server_name_rejects_malformed_names(monkeypatch, name): + monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_server_name", name) + with pytest.raises(ValueError, match="FLEXKV_RADIXSHMEM_SERVER_NAME"): + bootstrap.radix_server_name() # ============================================================================= @@ -1047,7 +986,7 @@ def test_prefetch_starts_a_peer_pull(): (call,) = engine.prefetch_calls # type: ignore[attr-defined] assert call["component_mask"] == _engine_mod.COMPONENT_MASK_FULL assert call["query_end"] == 4 - assert call["timeout_ms"] == load_radixshmem_config(None).client.prefetch_timeout_ms + assert call["timeout_ms"] == bootstrap.PREFETCH_TIMEOUT_MS def test_prefetch_without_peers_is_an_empty_plan(): @@ -1063,7 +1002,7 @@ def test_prefetch_without_peers_is_an_empty_plan(): def test_prefetch_backpressure_skips_the_peer_walk(): """Too many pulls in flight: no pull_async, so the client never blocks.""" engine = _global_cache_engine() - limit = load_radixshmem_config(None).client.prefetch_max_inflight + limit = bootstrap.PREFETCH_MAX_INFLIGHT engine._prefetch_jobs = [FakeJob(0, 4) for _ in range(limit)] # none done _graph, ops, return_mask = _run_get( engine, 4, _local_match([]), prefetch=True, prefetch_job=FakeJob(0, 4)) @@ -1354,10 +1293,11 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, from flexkv.common.config import CacheConfig, ModelConfig, SWAPoolConfig - rcfg = _radix_config(name=f"/swaplanner{os.getpid()}") - saved = {"enable_radixshmem": GLOBAL_CONFIG_FROM_ENV.enable_radixshmem} + server_name = f"/swaplanner{os.getpid()}" + saved = {"enable_radixshmem": GLOBAL_CONFIG_FROM_ENV.enable_radixshmem, + "radixshmem_server_name": GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name} GLOBAL_CONFIG_FROM_ENV.enable_radixshmem = True - set_radixshmem_config(rcfg) + GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name = server_name server = None engine = None @@ -1379,7 +1319,7 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, # test's; the planner's client brings the geometry. geo = bootstrap.expected_geometry(model_config, cache_config) data_bytes = num_blocks * geo.full_slot_bytes + swa_slots * geo.swa_slot_bytes - cfg = shmradix.ServerConfig(name=rcfg.server_name, data_bytes=data_bytes, + cfg = shmradix.ServerConfig(name=server_name, data_bytes=data_bytes, swa_ratio=swa_slots * geo.swa_slot_bytes / data_bytes, prefault=False) _sweep_region(cfg.name, cfg.resolved_data_name) @@ -1395,7 +1335,6 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, server.close() for name, value in saved.items(): setattr(GLOBAL_CONFIG_FROM_ENV, name, value) - set_radixshmem_config(None) def _split_swa(ops_of_type): @@ -1733,24 +1672,26 @@ def _node_main(rank, prefix, cluster_id, registry, rdma_dev, ready, done, output local_head_blocks=0): """One node: a data-mode RadixServer (in-process) plus the FlexKV engine.""" try: - endpoint = f"unix:///dev/shm/{prefix.lstrip('/')}_r{rank}.sock" + # Two servers on one host: distinct names, so radixshmem derives + # distinct sockets (cluster names are per node; the geometry is shared). + name = f"{prefix}_r{rank}" data_name = f"{prefix}_data_r{rank}" - _sweep_region(prefix, data_name) + _sweep_region(name, data_name) cluster_kwargs = dict( expected_min_nodes=2, registry=registry, cluster_id=cluster_id, node_name=f"r{rank}", rpc_address="0.0.0.0", index_dev=rdma_dev, gid_idx=int(os.getenv("FLEXKV_TEST_RADIX_GID_IDX", "3")), bootstrap_timeout_sec=60, rht_slots_per_bucket=4) cfg = shmradix.ServerConfig( - name=prefix, data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, slot_align=4096, + name=name, data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, slot_align=4096, data_name=data_name, prefault=False, transfer_devices=[rdma_dev], - endpoint=endpoint, cluster=shmradix.ClusterConfig(**cluster_kwargs), + cluster=shmradix.ClusterConfig(**cluster_kwargs), ) server = shmradix.RadixServer(cfg).start() # waiting: the engine's geometry starts the rendezvous - set_radixshmem_config(_radix_config(name=prefix, endpoint=endpoint, ready_timeout_s=180.0)) + bootstrap.READY_TIMEOUT_S = 180.0 # this process only: the rendezvous may take a while engine = CacheEngineRadixShmem( - prefix, geometry=shmradix.Geometry(block_size=16, full_slot_bytes=PEER_SLOT_BYTES, - slot_align=4096), + name, geometry=shmradix.Geometry(block_size=16, full_slot_bytes=PEER_SLOT_BYTES, + slot_align=4096), num_total_blocks=PEER_BLOCKS, tokens_per_block=16, peer_enabled=True) if not engine.peer_enabled: raise RuntimeError("engine did not see a distributed region") From fb3367eee982de2b054808c095b92aab53b06eb1 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 15:15:39 +0800 Subject: [PATCH 15/21] radixshmem: retry the attach only while nothing answers; the cluster question brings the geometry attach_radix_client retried every error from the RadixClient constructor as "server not up yet" at DEBUG level for READY_TIMEOUT_S and then reported "no radix-server reachable", which pointed operators at the wrong cause when a live server had refused the attach (INTERNAL: server is closed) or FlexKV had called it wrongly (TypeError). Now only radixshmem's UNAVAILABLE RuntimeError (nothing at the socket) and an RPC TimeoutError are retried; anything else is logged with the label and the socket and raised as is. The timeout message after wait_ready reads the server's state at that moment (client.status()) instead of the constructor-time info. radix_server_is_distributed(model_config, cache_config, ...) attaches with FlexKV's geometry like every other attach, so a server nobody has configured yet is configured by the question and answers from world_size instead of timing out with "nobody handed it a geometry". The sglang connector passes its model and cache config; the call sits after the SWA layer groups are set, so the geometry matches the KVManager's. Tests: attach errors other than UNAVAILABLE come back at once, UNAVAILABLE is retried until the deadline; an unconfigured server answers the cluster question and a geometry-less attach on it times out naming the cause. Verified: 444 CPU tests; GPU e2e dp_size 1 and 2 pass. --- docs/radixshmem/config_zh.md | 5 +- flexkv/integration/sglang/connector.py | 11 +++-- flexkv/server/shm_radix_bootstrap.py | 43 +++++++++++++--- tests/radixshmem/test_radix_shmem_engine.py | 55 ++++++++++++++++++++- 4 files changed, 100 insertions(+), 14 deletions(-) diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 8dca2491d..3d41f983b 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -132,5 +132,6 @@ FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTA - `FLEXKV_RADIXSHMEM_SERVER_NAME` 须以 `/` 开头且无空白。 - `CacheConfig`:`FLEXKV_CPU_LAYOUT != BLOCKFIRST`;打开了 `enable_ssd`、`enable_remote`、`enable_p2p_cpu` 或 `enable_p2p_ssd`; SWA 开启但 `window_blocks < 1`。 -- attach:600 s 内连不上 server 报 `no radix-server named ... reachable ...`(附启动命令);server 一直在等几何或 - 配置失败报 `not ready within ...`(附 server 的 `mode` 和 `last_error`);几何冲突见第 3 节。 +- attach:600 s 内连不上 server 报 `no radix-server named ... reachable ...`(附启动命令);server 有响应但 attach + 失败(如 `server is closed`)立即报错;server 一直在等几何或配置失败报 `not ready within ...`(附 server 当时的 + `mode` 和 `last_error`);几何冲突见第 3 节。 diff --git a/flexkv/integration/sglang/connector.py b/flexkv/integration/sglang/connector.py index 18b938b5b..40da36e4e 100644 --- a/flexkv/integration/sglang/connector.py +++ b/flexkv/integration/sglang/connector.py @@ -62,13 +62,15 @@ from flexkv.transfer_manager import TransferManagerOnRemote -def _radixshmem_distributed() -> bool: +def _radixshmem_distributed(model_config, cache_config) -> bool: """Whether this node's radix-server is part of a cluster (peer pulls possible). The server is the operator's process and knows; every TP rank asks it the same question once it is ready, so the prefetch gate below is - the same in all ranks (the PREFETCH_START scatter needs that).""" + the same in all ranks (the PREFETCH_START scatter needs that). The attach + brings FlexKV's geometry, so the answer does not depend on which process + configured the server first.""" from flexkv.server.shm_radix_bootstrap import radix_server_is_distributed - return radix_server_is_distributed(label="FlexKVConnector") + return radix_server_is_distributed(model_config, cache_config, label="FlexKVConnector") logger = logging.getLogger(__name__) @@ -341,7 +343,8 @@ def __init__( or self.cache_config.enable_kv_sharing # radixshmem cluster: prefetch is where a peer's blocks are pulled # into this node (RadixClient.pull_async); GET then matches locally. - or (GLOBAL_CONFIG_FROM_ENV.enable_radixshmem and _radixshmem_distributed()) + or (GLOBAL_CONFIG_FROM_ENV.enable_radixshmem + and _radixshmem_distributed(self.model_config, self.cache_config)) ) self._shutdown_done = False diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 58b5d62e0..0e169bdfd 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -266,6 +266,25 @@ def expected_geometry(model_config: ModelConfig, cache_config: CacheConfig) -> R # -------------------------------------------------------------------- attach +def _nothing_answers(e: BaseException) -> bool: + """Whether an attach error means no server answered at the socket yet. + radixshmem's RPC layer raises ``RuntimeError("UNAVAILABLE: ...")`` for a + missing or unreachable server and ``TimeoutError`` for an RPC deadline; + both are worth retrying. Anything else is a live server's answer (e.g. + ``INTERNAL: server is closed``) or a programming error and is raised as is.""" + return isinstance(e, TimeoutError) or ( + isinstance(e, RuntimeError) and str(e).startswith("UNAVAILABLE")) + + +def _current_status(client: "shmradix.RadixClient"): + """The server's state now (one RPC); the constructor-time info when the + server cannot be asked any more.""" + try: + return client.status() + except Exception: # noqa: BLE001 - gone or unreachable: report what we had + return client.info + + def attach_radix_client(name: Optional[str] = None, *, geometry: Any = None, @@ -312,7 +331,13 @@ def attach_radix_client(name: Optional[str] = None, raise ValueError( f"{label}: radix-server {name} cannot serve FlexKV's geometry ({e}); check its " f"--data-bytes / --swa-ratio") from e - except Exception as e: # noqa: BLE001 - not reachable yet: no socket, no listener + except Exception as e: # noqa: BLE001 - classified below + if not _nothing_answers(e): + # A live server refused the attach, or FlexKV called it wrongly: + # retrying would only hide the cause for READY_TIMEOUT_S. + flexkv_logger.error( + f"{label}: attach to radix-server {name} at {where} failed: {e!r}") + raise last = e if time.monotonic() >= deadline: raise TimeoutError( @@ -326,7 +351,8 @@ def attach_radix_client(name: Optional[str] = None, try: info = client.wait_ready(remaining) except TimeoutError as e: - mode, err = client.info.mode, client.info.last_error + info = _current_status(client) + mode, err = info.mode, info.last_error client.close() raise TimeoutError( f"{label}: radix-server {name} not ready within {timeout_s:.0f}s (mode={mode}" @@ -488,16 +514,19 @@ def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, client.close() -def radix_server_is_distributed(name: Optional[str] = None, +def radix_server_is_distributed(model_config: ModelConfig, cache_config: CacheConfig, *, + name: Optional[str] = None, timeout_s: Optional[float] = None, label: str = "radixshmem") -> bool: """Whether this node's radix-server is part of a cluster (world_size > 1), i.e. whether peer pulls are possible. Asked of the server itself once it - is ready (a geometry-less attach that only waits), so every process that - asks gets the same answer -- the framework adapters gate the prefetch path - on it in every TP rank, and the ranks must agree.""" - client = attach_radix_client(name, timeout_s=timeout_s, label=label) + is ready, so every process that asks gets the same answer -- the framework + adapters gate the prefetch path on it in every TP rank, and the ranks must + agree. The attach brings FlexKV's geometry like every other one, so the + question can be asked of a server nobody has configured yet.""" + geometry = expected_geometry(model_config, cache_config) + client = attach_radix_client(name, geometry=geometry, timeout_s=timeout_s, label=label) try: return int(client.info.world_size) > 1 finally: diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 556c76020..945bf3cbe 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -682,7 +682,8 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): client, dataclasses.replace(geo, register_chunk_tokens=4096), "test") assert bootstrap.radix_cluster_rank(client) == 0 and client.info.world_size == 1 # the connectors' prefetch gate asks the server the same question - assert bootstrap.radix_server_is_distributed(name, timeout_s=30) is False + assert bootstrap.radix_server_is_distributed( + model_config, cache_config, name=name, timeout_s=30) is False # another expectation against the same regions fails closed cache_config.tokens_per_block = 32 with pytest.raises(ValueError, match="tokens_per_block"): @@ -725,6 +726,58 @@ def _start_later(): bootstrap.attach_radix_client(f"/nobody{os.getpid()}", geometry=geo, timeout_s=2) +def test_attach_retries_only_while_nothing_answers(env, monkeypatch): + """An unreachable socket (UNAVAILABLE) is retried until the deadline. Any + other error is a live server's answer or a bug and comes back at once, + instead of being retried for READY_TIMEOUT_S and reported as 'no server'.""" + name = f"/errs{os.getpid()}" + calls = [] + + def _ctor_raising(exc): + def _ctor(*args, **kwargs): + calls.append(exc) + raise exc + return _ctor + + for exc, exc_type, match in ( + (TypeError("unexpected keyword argument 'endpoint'"), TypeError, "unexpected keyword"), + (RuntimeError("INTERNAL: server is closed"), RuntimeError, "server is closed")): + calls.clear() + monkeypatch.setattr(shmradix, "RadixClient", _ctor_raising(exc)) + t0 = time.monotonic() + with pytest.raises(exc_type, match=match): + bootstrap.attach_radix_client(name, timeout_s=30) + assert len(calls) == 1 and time.monotonic() - t0 < 5 + calls.clear() + monkeypatch.setattr(shmradix, "RadixClient", + _ctor_raising(RuntimeError("UNAVAILABLE: failed to connect to all addresses"))) + with pytest.raises(TimeoutError, match="no radix-server named"): + bootstrap.attach_radix_client(name, timeout_s=1.5) + assert len(calls) >= 2 # retried until the deadline + + +def test_unconfigured_server_answers_the_cluster_question(env): + """`radix_server_is_distributed` brings the geometry, so a server nobody + configured yet answers instead of timing out; a geometry-less attach on + that server names the cause, read from the server's state at that moment.""" + model_config, cache_config = _configs(num_cpu_blocks=64) + geo = bootstrap.expected_geometry(model_config, cache_config) + name = f"/waiting{os.getpid()}" + env.server(shmradix.ServerConfig(name=name, data_bytes=64 * geo.full_slot_bytes, + prefault=False)) + with pytest.raises(TimeoutError, match="mode=waiting; nobody handed it a geometry"): + bootstrap.attach_radix_client(name, timeout_s=2) + t0 = time.monotonic() + assert bootstrap.radix_server_is_distributed( + model_config, cache_config, name=name, timeout_s=60) is False + assert time.monotonic() - t0 < 30 + client = bootstrap.attach_radix_client(name, geometry=geo, timeout_s=60) + try: + assert client.info.mode == "ready" and int(client.mempool_total()) == 64 + finally: + client.close() + + # ============================================================================= # Part 1c — FLEXKV_RADIXSHMEM_SERVER_NAME, the one radixshmem-mode setting # (shm_radix_bootstrap.radix_server_name); everything else about the attach From 94f6381fa9adf1e49cd32d6c78fa4f17a0f8a0e8 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 15:28:43 +0800 Subject: [PATCH 16/21] radixshmem: say at start-up that cpu_cache_gb does not size the CPU tier In radixshmem mode the CPU tier is the radix-server's SlotStore, sized by its --data-bytes (and --swa-ratio); the slot counts it plans replace num_cpu_blocks / swa.num_slots. KVManager now logs that when it starts: shm_radix_bootstrap.cpu_sizing_notice() gives the level and the text, a WARNING with the configured value when the deployment set cpu_cache_gb (CacheConfig._user_cpu_cache_gb), INFO for a bare default. Unit test for the notice; the doc's overview mentions the warning. --- docs/radixshmem/config_zh.md | 2 +- flexkv/kvmanager.py | 6 ++++-- flexkv/server/shm_radix_bootstrap.py | 22 ++++++++++++++++++++- tests/radixshmem/test_radix_shmem_engine.py | 15 ++++++++++++++ 4 files changed, 41 insertions(+), 4 deletions(-) diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 3d41f983b..e850c771f 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -10,7 +10,7 @@ FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取) FlexKV 只实例化 `shmradix.RadixClient`:第一个 client 把几何(每 block 的 token 数、一个 CPU block 和一个 SWA page 的字节数、SWA 窗口、slot 对齐)交给 server,server 按 `--data-bytes` 和 `--swa-ratio` 规划各池的 slot 数并发布, 每个 FlexKV 进程 attach 时把 slot 数采纳到 `CacheConfig`(`num_cpu_blocks`、`swa.num_slots`)。CPU 层容量由 -`radix-server --data-bytes` 决定,`cpu_cache_gb` 在该模式下不起作用。 +`radix-server --data-bytes` 决定,`cpu_cache_gb` 在该模式下不起作用(KVManager 启动时打印一条 WARNING 提示)。 实现:`flexkv/server/shm_radix_bootstrap.py`(几何、attach、采纳、固定参数)。 radixshmem 侧接口见 radixshmem 仓库 `python/README.md`。 diff --git a/flexkv/kvmanager.py b/flexkv/kvmanager.py index e103c6349..76a496def 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -92,8 +92,10 @@ def __init__(self, ) if self.enable_radixshmem: - flexkv_logger.info("[KVManager] radixshmem mode: the CPU tier is the " - "operator's radix-server") + # Say up front that the CPU tier is not sized by cpu_cache_gb here. + from flexkv.server.shm_radix_bootstrap import cpu_sizing_notice + level, text = cpu_sizing_notice(cache_config, GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name) + getattr(flexkv_logger, level)(f"[KVManager] {text}") # Multi-instance mode also requires server_client_mode self.server_client_mode = (model_config.dp_size > 1 or diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 0e169bdfd..83dd50682 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -34,7 +34,7 @@ import dataclasses import time -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple import torch @@ -80,6 +80,26 @@ def radix_server_name() -> str: return name +def cpu_sizing_notice(cache_config: CacheConfig, server_name: str) -> Tuple[str, str]: + """(log level, text) telling the operator at start-up that ``cpu_cache_gb`` + does not size the CPU tier in radixshmem mode: the radix-server's budget + does, and its slot counts replace ``num_cpu_blocks`` / ``swa.num_slots``. + A warning when the deployment configured cpu_cache_gb itself (it expects + the value to matter), info for a bare default.""" + # CacheConfig keeps the deployment's cpu_cache_gb only as _user_cpu_cache_gb + # (0 when it was never given); num_cpu_blocks is what it was turned into. + user_gb = float(getattr(cache_config, "_user_cpu_cache_gb", 0) or 0) + setting = f"cpu_cache_gb={user_gb:g}" if user_gb > 0 else "cpu_cache_gb" + swa = cache_config.swa + has_swa = swa is not None and swa.enabled + replaced = "num_cpu_blocks" + (" and swa.num_slots" if has_swa else "") + text = (f"radixshmem mode: {setting} is ignored. The CPU tier is radix-server {server_name}'s " + f"SlotStore, sized by its --data-bytes" + (" and --swa-ratio" if has_swa else "") + + f"; the slot counts it planned replace {replaced} " + f"(num_cpu_blocks={cache_config.num_cpu_blocks} until then)") + return ("warning" if user_gb > 0 else "info"), text + + def default_endpoint(server_name: str) -> str: """The gRPC socket radixshmem derives from ``radix-server --name ``: ``unix:///dev/shm/ '_'>.sock``.""" diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 945bf3cbe..73f3adc4d 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -797,6 +797,21 @@ def test_radix_server_name_follows_the_env(monkeypatch): assert 0 < bootstrap.PREFETCH_MAX_INFLIGHT < bootstrap.MAX_OUTSTANDING +def test_cpu_sizing_notice_names_the_ignored_setting(): + """Start-up tells the operator that cpu_cache_gb does not size the CPU tier + in this mode: a warning when the deployment set it, info for the default.""" + model_config, cache_config = _configs(num_cpu_blocks=64, swa_slots=16) + cache_config._user_cpu_cache_gb = 32 + level, text = bootstrap.cpu_sizing_notice(cache_config, "/flexkv") + assert level == "warning" + assert "cpu_cache_gb=32 is ignored" in text + assert "radix-server /flexkv" in text and "--data-bytes and --swa-ratio" in text + assert "num_cpu_blocks and swa.num_slots (num_cpu_blocks=64 until then)" in text + _, cache_config = _configs(num_cpu_blocks=64) + level, text = bootstrap.cpu_sizing_notice(cache_config, "/flexkv") + assert level == "info" and "cpu_cache_gb is ignored" in text and "--swa-ratio" not in text + + @pytest.mark.parametrize("name", ["kv", "/a b", "/"]) def test_radix_server_name_rejects_malformed_names(monkeypatch, name): monkeypatch.setattr(GLOBAL_CONFIG_FROM_ENV, "radixshmem_server_name", name) From 44d38b1a53f412c02520b68e0ece7d9528bd2e3e Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 16:44:34 +0800 Subject: [PATCH 17/21] radixshmem: one --transfer-dev with every HCA in the cluster example The multi-node radix-server example lists mlx5_0..mlx5_7 in one comma-separated --transfer-dev (radixshmem MR !18 adds the comma form; repeating the flag still works). Co-Authored-By: Claude Fable 5.1 --- docs/radixshmem/config_zh.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index e850c771f..1e80e54f1 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -93,7 +93,7 @@ radix-server --name /flexkv --data-bytes 64G --swa-ratio 0.5 \ --expected-min-nodes 4 --num-rht-shards 4 --rht-slots 4 \ --registry etcd://10.0.0.1:2379 --cluster-id prod_a \ --rpc-interface bond0 --index-dev mlx5_bond_0 --gid-idx 3 \ - --transfer-dev mlx5_1 --transfer-dev mlx5_2 --bootstrap-timeout 600 + --transfer-dev mlx5_0,mlx5_1,mlx5_2,mlx5_3,mlx5_4,mlx5_5,mlx5_6,mlx5_7 --bootstrap-timeout 600 ``` 集群一致的几何字段(`block_size`、池集合、每池 `slot_bytes`、SWA 窗口、`slot_align`、`register_chunk_tokens`)由第一个 From 4e736dee72626f749647ab1b9b3081ccd4db3d56 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 20:46:08 +0800 Subject: [PATCH 18/21] radixshmem: attach the index at start-up and retry a peer RHT shard that is not up yet attach_radix_client(attach_index=True) brings client.index up before returning instead of on first use. On a cluster that opens the RDMA queue pairs to every peer's RHT shard holder; right after the rendezvous a holder may not accept yet and radixshmem gives up on the connect with "RhtConsumer: failed to connect RHT holder". That error is now retried every 5 s until the ready deadline; any other error is raised at once. CacheEngineRadixShmem and the TE's TransferManager turn it on; callers that only read info (adopting the slot counts, asking whether the server is distributed) leave it off and open no RDMA state. Seen on the two-node cp4x4 bed: node 24's engine attached about 60 s after "Cluster ready (2 nodes)" and failed the XRC connect to node 25's holder while four FlexKV processes attached at once. With the retry the attach went through after 7-8 rounds; a later run needed none. --- flexkv/cache/radix_shmem_engine.py | 2 +- flexkv/server/shm_radix_bootstrap.py | 32 +++++++++++++++ flexkv/transfer_manager.py | 3 +- tests/radixshmem/test_radix_shmem_engine.py | 45 +++++++++++++++++++++ 4 files changed, 80 insertions(+), 2 deletions(-) diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 8019053b2..1b502cfc8 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -171,7 +171,7 @@ def __init__(self, cpu_swa = swa_config.for_cache_tier(DeviceType.CPU) if swa_config is not None else None self.swa_enabled = cpu_swa is not None and cpu_swa.num_slots > 0 - self._client = attach_radix_client(server_name, geometry=geometry, + self._client = attach_radix_client(server_name, geometry=geometry, attach_index=True, label="CacheEngineRadixShmem") self._tree = self._client # index ops pass through the client self.shm_name = self._client.info.index_name # node-suffixed when distributed diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 83dd50682..5211fe600 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -310,10 +310,16 @@ def attach_radix_client(name: Optional[str] = None, geometry: Any = None, timeout_s: Optional[float] = None, max_outstanding: Optional[int] = None, + attach_index: bool = False, label: str = "radixshmem") -> "shmradix.RadixClient": """A ready ``shmradix.RadixClient`` on the radix-server ``name`` (default: :func:`radix_server_name`), at the socket radixshmem derives from the name. + ``attach_index=True`` also brings the index client up before returning + (RDMA queue pairs to every peer's RHT shard on a cluster) and retries a + peer that does not accept yet; callers that only read ``info`` (adopting + counts, asking about the cluster) leave it off and open no RDMA state. + With ``geometry`` (a :class:`RadixGeometry` or a ``shmradix.Geometry``) the client hands the server FlexKV's slot shape on the way; the server plans the counts from its budget. That is idempotent, so every FlexKV process @@ -382,6 +388,8 @@ def attach_radix_client(name: Optional[str] = None, except RuntimeError as e: client.close() raise RuntimeError(f"{label}: radix-server {name} failed to configure: {e}") from e + if attach_index: + _attach_index(client, name, deadline, label) flexkv_logger.info( f"{label}: attached radix-server {name} ({where}): index={info.index_name}, " f"rank={info.rank}/{info.world_size}, data_plane={info.data_plane}, " @@ -389,6 +397,30 @@ def attach_radix_client(name: Optional[str] = None, return client +def _attach_index(client: "shmradix.RadixClient", name: str, deadline: float, label: str) -> None: + """Bring the index client up now (``client.index``) instead of on first use. + On a cluster this opens the RDMA queue pairs to every peer's RHT shard; + right after the rendezvous a peer's holder may not accept yet, and + radixshmem gives up on such a connect with ``RhtConsumer: failed to + connect RHT holder``. That is retried until ``deadline``; any other error + is raised at once.""" + attempt = 0 + while True: + try: + client.index + return + except RuntimeError as e: + transient = "RHT" in str(e) or "RhtConsumer" in str(e) + if not transient or time.monotonic() >= deadline: + client.close() + raise RuntimeError(f"{label}: radix-server {name}: attaching the index failed: {e}") from e + attempt += 1 + flexkv_logger.warning( + f"{label}: radix-server {name}: RHT peer not reachable yet ({e}); retry {attempt} " + f"in 5s") + time.sleep(5.0) + + def _describe_published(g: Optional[Dict[str, Any]]) -> str: if not g: return "(none)" diff --git a/flexkv/transfer_manager.py b/flexkv/transfer_manager.py index f312ed9c1..d2d6dd7ab 100644 --- a/flexkv/transfer_manager.py +++ b/flexkv/transfer_manager.py @@ -439,7 +439,8 @@ def initialize_transfer_engine(self) -> None: from flexkv.server.shm_radix_bootstrap import (adopt_geometry, attach_radix_client, check_geometry, expected_geometry) geometry = expected_geometry(self.model_config, self.cache_config) - radix_client = attach_radix_client(geometry=geometry, label="TransferManager") + radix_client = attach_radix_client(geometry=geometry, attach_index=True, + label="TransferManager") check_geometry(radix_client, geometry, label="TransferManager") adopt_geometry(self.cache_config, radix_client, label="TransferManager") self._radix_client = radix_client diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 73f3adc4d..9e2311228 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -756,6 +756,51 @@ def _ctor(*args, **kwargs): assert len(calls) >= 2 # retried until the deadline +def test_attach_retries_a_peer_rht_shard_that_is_not_up_yet(env, monkeypatch): + """Right after a cluster rendezvous a peer's RHT holder may refuse the RDMA + connect; radixshmem raises 'RhtConsumer: failed to connect RHT holder' from + the index attach. attach_radix_client retries that until the deadline and + still raises other index errors at once.""" + class _Fake: + def __init__(self, failures, error="RhtConsumer: failed to connect RHT holder"): + self.failures, self.error, self.closed = failures, error, False + self.info = SimpleNamespace(mode="ready", index_name="/x", rank=0, world_size=2, + data_plane=True, geometry=None, last_error="") + self.name = "/x" + + def wait_ready(self, timeout_s): + return self.info + + @property + def index(self): + if self.failures > 0: + self.failures -= 1 + raise RuntimeError(self.error) + return object() + + def close(self): + self.closed = True + + fakes = [] + + def _ctor(name, spec, max_outstanding=256): + fakes.append(_Fake(fakes and fakes[-1].failures or 2)) # first client: 2 failures, then ok + return fakes[-1] + + monkeypatch.setattr(bootstrap.time, "sleep", lambda s: None) + monkeypatch.setattr(shmradix, "RadixClient", _ctor) + client = bootstrap.attach_radix_client("/x", timeout_s=60, attach_index=True) + assert client is fakes[0] and fakes[0].failures == 0 and not fakes[0].closed + fakes.clear() + monkeypatch.setattr(shmradix, "RadixClient", lambda *a, **k: _Fake(1, "index/store/geometry mismatch: x")) + with pytest.raises(RuntimeError, match="attaching the index failed"): + bootstrap.attach_radix_client("/x", timeout_s=60, attach_index=True) + # without attach_index the index is left alone (info-only callers open no RDMA state) + fakes.clear() + monkeypatch.setattr(shmradix, "RadixClient", lambda *a, **k: _Fake(5)) + assert bootstrap.attach_radix_client("/x", timeout_s=60).failures == 5 + + def test_unconfigured_server_answers_the_cluster_question(env): """`radix_server_is_distributed` brings the geometry, so a server nobody configured yet answers instead of timing out; a geometry-less attach on From d952491f5d3904cc24324ea752377d049288c5b3 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Tue, 22 Sep 2026 23:00:05 +0800 Subject: [PATCH 19/21] mooncake store: use the public batch_is_exist and read put()'s per-key result MooncakeStoreClient.exists() called the SDK's private _batch_exist, which mooncake-transfer-engine 0.3.12 no longer has (batch_is_exist is the public name, and batch_exists_impl already used it). The single-key put() compared the list batch_put_from returns with 0 and so reported every successful put as a failure; check the one element like batch_put does. Neither is on FlexKV's transfer hot path (that uses batch_put / batch_get / batch_exists), found while bringing up the shared Mooncake Store arm of the cp4x4 comparison (4 sglang instances, DSv4-Flash, two nodes). --- flexkv/external/mooncake_store_utils.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/flexkv/external/mooncake_store_utils.py b/flexkv/external/mooncake_store_utils.py index d2b9da122..ca453d348 100644 --- a/flexkv/external/mooncake_store_utils.py +++ b/flexkv/external/mooncake_store_utils.py @@ -342,7 +342,7 @@ def put(self, key: str, buffer_ptr: int, buffer_size: int) -> bool: return True ret_code = self._store.batch_put_from([key], [buffer_ptr], [buffer_size]) - return ret_code == 0 + return ret_code[0] == 0 def batch_put( self, @@ -421,7 +421,7 @@ def batch_exists(self, keys_strs: list[str]) -> int: def exists(self, key: str) -> bool: """Check existence of a key in the store.""" self._ensure_setup() - result = self._store._batch_exist([key]) + result = self._store.batch_is_exist([key]) return result[0] == 1 def zero_copy_put_impl( From 1f0d88f596ed71817d6e3a1c73d9696f43ac3ca2 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Mon, 28 Sep 2026 19:04:25 +0800 Subject: [PATCH 20/21] radixshmem: a failed cluster-wide flush no longer recycles slots the tree owns CacheEngineRadixShmem.insert() calls _tree.flush() after _tree.insert() has taken the slots (auto_recycle=True). When the flush raised, e.g. with "DcInitiator::flush WC error status=10 vendor=136" from an RDMA write to a peer's RHT shard, StagedRadixInsert.publish() caught the exception and recycled the same slots into the mempool. The slots then sat in the tree and on the free list at once, so a later PUT could overwrite KV that the tree still served for the old prefix. The flush only makes the blocks routable from other nodes; the local insert has already succeeded. insert() now logs the flush failure as a warning and returns normally, so publish() recycles only when the insert itself fails. test_failed_cluster_publish_keeps_the_slots_in_the_tree makes flush raise and checks that the inserted slots stay matchable and are not handed out again by take(). It fails without this change. Co-Authored-By: Claude Opus 5.5 (1M context) --- flexkv/cache/radix_shmem_engine.py | 9 +++++- tests/radixshmem/test_radix_shmem_engine.py | 33 +++++++++++++++++++++ 2 files changed, 41 insertions(+), 1 deletion(-) diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 1b502cfc8..5cdfa8889 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -342,7 +342,14 @@ def insert(self, return if self.peer_enabled and component == COMPONENT_FULL: - self._tree.flush() # make the new blocks visible cluster-wide + # The tree owns the slots by now: raising would make the caller + # recycle them a second time. + try: + self._tree.flush() # make the new blocks visible cluster-wide + except Exception as e: + flexkv_logger.warning( + f"radixshmem insert on {self.shm_name}: {landed} blocks landed " + f"locally but the cluster-wide publish failed: {e}") if (self.event_collector is not None and component == COMPONENT_FULL and result.error == shmradix.InsertError.OK): diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index 9e2311228..b35c6f12e 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -98,6 +98,7 @@ def _load_module_direct(name: str, path: str): # The planner only duck-types the match, so the side-loaded class is as good as # the one `flexkv.cache.cache_engine` imports — and it needs no c_ext. ShmRadixMatch = _engine_mod.ShmRadixMatch +StagedRadixInsert = _engine_mod.StagedRadixInsert # Pure-Python (no c_ext): the bootstrap (server config, geometry, attach) and # the transfer enums. @@ -296,6 +297,38 @@ def test_insert_publishes_immediately(env): r.release() +def test_failed_cluster_publish_keeps_the_slots_in_the_tree(env): + """A flush error (e.g. an RDMA write to a peer's RHT shard failing) comes + after the tree took the slots: publish() must not recycle them again, or + the mempool hands out slots the tree still serves.""" + engine, _server = env.make(f"/cers_flushfail{os.getpid()}", blocks=64) + + class _FlushFails: + def __init__(self, tree): + self._inner = tree + + def __getattr__(self, name): + return getattr(self._inner, name) + + def flush(self): + raise RuntimeError("DcInitiator::flush WC error status=10 vendor=136") + + engine._tree = _FlushFails(engine._tree) + engine.peer_enabled = True + + seq = FakeSeq(block_hashes=_hashes(seed=9, num=6)) + slots = engine.take(num_required_blocks=6) + StagedRadixInsert(engine, seq, slots, path_end=6, label="PUT").publish() + + r = engine.match(seq) + assert r.num_matched_blocks == 6 + np.testing.assert_array_equal(np.sort(r.local_slots), np.sort(slots)) + rest = engine.take(num_required_blocks=64) # the pin keeps them from eviction + assert not np.isin(slots, rest).any() + engine.recycle(rest) + r.release() + + def test_recycle_returns_staged_slots(env): """Slots whose transfer never landed are only reachable through recycle(). From 67302abb48b3b6ef6bc0c6dbe53846c2cdc2d190 Mon Sep 17 00:00:00 2001 From: Zhuofan Li Date: Fri, 9 Oct 2026 16:26:19 +0800 Subject: [PATCH 21/21] radixshmem: follow the shmradix -> radixshmem package rename radixshmem on GitHub (ai-dynamo/radixshmem, 3608a23) renamed its Python package from shmradix to radixshmem; the API is otherwise unchanged. Import radixshmem everywhere, rename RadixGeometry.to_shmradix / _ensure_shmradix to match, start the test server with python -m radixshmem.cli, and require 3608a23 or later in the changelog. --- CHANGELOG.md | 2 +- docs/radixshmem/config_zh.md | 4 +- flexkv/cache/radix_shmem_engine.py | 32 ++++---- flexkv/kvtask.py | 2 +- flexkv/server/server.py | 2 +- flexkv/server/shm_radix_bootstrap.py | 54 ++++++------- flexkv/storage/allocator.py | 4 +- flexkv/storage/storage_engine.py | 4 +- tests/radixshmem/radix_e2e_common.py | 4 +- .../radixshmem/test_e2e_radix_prefetch_p2p.py | 4 +- tests/radixshmem/test_e2e_radix_shmem.py | 2 +- tests/radixshmem/test_radix_shmem_engine.py | 76 +++++++++---------- 12 files changed, 95 insertions(+), 95 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5875d3e55..32865c913 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,7 +10,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Feature Universal: -- Add radixshmem mode (`FLEXKV_ENABLE_RADIXSHMEM=1`; requires radixshmem f939910 or later): the CPU tier is a per-node `radix-server` run by the operator (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`) that owns the radix index, the SlotStore and the cross-node RDMA pull. FlexKV attaches with `shmradix.RadixClient(name, Geometry)`: it hands the server its slot geometry (tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), adopts the slot counts the server plans from its budget into `CacheConfig` (`shm_radix_bootstrap.adopt_radix_server`; `cpu_cache_gb` has no effect in this mode) and uses the server's SlotStore as the CPU pool in the TE and every transfer worker; cross-node reuse is `RadixClient.pull_async` from the prefetch path whenever the server runs as a cluster. `FLEXKV_RADIXSHMEM_SERVER_NAME` (default `/flexkv`) names the server to attach to; the endpoint is the one radixshmem derives from the name, and the ready wait (600 s) and the prefetch limits are fixed defaults. Process model as in the other modes: engine mode for one DP, the node's KVServer for dp_size > 1 or several `FLEXKV_INSTANCE_NUM` engines. Planning is `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, a `GlobalCacheEngine` subclass over the `_prepare_request` / `_build_cpu_cache_engine` hooks) on `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`); a plan cancelled before launch or failing during planning returns its slots and match pin. CPU tier only (`enable_ssd`, `enable_remote`, `enable_p2p_*` must be off). Reference: `docs/radixshmem/config_zh.md`, launch scripts in `examples/radixshmem/`. +- Add radixshmem mode (`FLEXKV_ENABLE_RADIXSHMEM=1`; requires [radixshmem](https://github.com/ai-dynamo/radixshmem) 3608a23 or later): the CPU tier is a per-node `radix-server` run by the operator (`radix-server --name /flexkv --data-bytes 64G [--swa-ratio R] [cluster flags]`) that owns the radix index, the SlotStore and the cross-node RDMA pull. FlexKV attaches with `radixshmem.RadixClient(name, Geometry)`: it hands the server its slot geometry (tokens per block, bytes per CPU block / SWA page, SWA window, slot alignment), adopts the slot counts the server plans from its budget into `CacheConfig` (`shm_radix_bootstrap.adopt_radix_server`; `cpu_cache_gb` has no effect in this mode) and uses the server's SlotStore as the CPU pool in the TE and every transfer worker; cross-node reuse is `RadixClient.pull_async` from the prefetch path whenever the server runs as a cluster. `FLEXKV_RADIXSHMEM_SERVER_NAME` (default `/flexkv`) names the server to attach to; the endpoint is the one radixshmem derives from the name, and the ready wait (600 s) and the prefetch limits are fixed defaults. Process model as in the other modes: engine mode for one DP, the node's KVServer for dp_size > 1 or several `FLEXKV_INSTANCE_NUM` engines. Planning is `RadixShmemCacheEngine` (`flexkv/cache/radix_shmem_planner.py`, a `GlobalCacheEngine` subclass over the `_prepare_request` / `_build_cpu_cache_engine` hooks) on `CacheEngineRadixShmem` (`flexkv/cache/radix_shmem_engine.py`); a plan cancelled before launch or failing during planning returns its slots and match pin. CPU tier only (`enable_ssd`, `enable_remote`, `enable_p2p_*` must be off). Reference: `docs/radixshmem/config_zh.md`, launch scripts in `examples/radixshmem/`. - `KVServer.create_server(inherit_env=False)` passes `PYTHONPATH`, `LD_LIBRARY_PATH` and `PATH` to the server child along with the `FLEXKV_*` variables. - `gen_hashes` / `Hasher.update` hash numpy buffers directly (`c_ext.gen_hashes_numpy` / `update_numpy`, safe to call concurrently); `gen_hashes_numpy` validates dtype (int64 tokens, uint64 hashes), contiguity and sizes. diff --git a/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md index 1e80e54f1..e7025410d 100644 --- a/docs/radixshmem/config_zh.md +++ b/docs/radixshmem/config_zh.md @@ -7,7 +7,7 @@ FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取) | 运维 | `radix-server` 命令行,每节点一个进程 | 名字、SlotStore 字节预算与 SWA 占比、hugepage、传输引擎、集群成员(etcd、网卡、rank)、索引调优 | | FlexKV | 两个环境变量 | 是否启用、attach 哪个 server | -FlexKV 只实例化 `shmradix.RadixClient`:第一个 client 把几何(每 block 的 token 数、一个 CPU block 和一个 SWA page +FlexKV 只实例化 `radixshmem.RadixClient`:第一个 client 把几何(每 block 的 token 数、一个 CPU block 和一个 SWA page 的字节数、SWA 窗口、slot 对齐)交给 server,server 按 `--data-bytes` 和 `--swa-ratio` 规划各池的 slot 数并发布, 每个 FlexKV 进程 attach 时把 slot 数采纳到 `CacheConfig`(`num_cpu_blocks`、`swa.num_slots`)。CPU 层容量由 `radix-server --data-bytes` 决定,`cpu_cache_gb` 在该模式下不起作用(KVManager 启动时打印一条 WARNING 提示)。 @@ -43,7 +43,7 @@ attach 的其余参数是 `flexkv/server/shm_radix_bootstrap.py` 里的常量: ## 3. 几何与 slot 数 -FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `shmradix.Geometry`): +FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `radixshmem.Geometry`): | 字段 | 来源 | |---|---| diff --git a/flexkv/cache/radix_shmem_engine.py b/flexkv/cache/radix_shmem_engine.py index 5cdfa8889..fafff44a3 100644 --- a/flexkv/cache/radix_shmem_engine.py +++ b/flexkv/cache/radix_shmem_engine.py @@ -1,7 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # cython: boundscheck=True, wraparound=True """ -The radixshmem CPU tier: one process's `shmradix.RadixClient` on the node's +The radixshmem CPU tier: one process's `radixshmem.RadixClient` on the node's radix-server (the operator's process: index shm + SlotStore + RDMA engine; the attach and the geometry hand-off are in `flexkv.server.shm_radix_bootstrap`). Used only by `flexkv.cache.radix_shmem_planner.RadixShmemCacheEngine`. @@ -36,16 +36,16 @@ from flexkv.integration.dynamo.collector import KVEventCollector try: - import shmradix + import radixshmem except ImportError as e: # pragma: no cover raise ImportError( - "shmradix is not installed; install it from the radixshmem repo " + "radixshmem is not installed; install it from the radixshmem repo " "(pip install -e radixshmem/python)") from e -COMPONENT_MASK_FULL = int(shmradix.COMPONENT_MASK_FULL) -COMPONENT_MASK_SWA = int(shmradix.COMPONENT_MASK_SWA) -COMPONENT_FULL = shmradix.ComponentType.FULL -COMPONENT_SWA = shmradix.ComponentType.SWA +COMPONENT_MASK_FULL = int(radixshmem.COMPONENT_MASK_FULL) +COMPONENT_MASK_SWA = int(radixshmem.COMPONENT_MASK_SWA) +COMPONENT_FULL = radixshmem.ComponentType.FULL +COMPONENT_SWA = radixshmem.ComponentType.SWA def _empty_i64() -> np.ndarray: @@ -93,7 +93,7 @@ def __init__(self, path_end: int, label: str, holds: Sequence[Callable[[], None]] = (), - component: shmradix.ComponentType = COMPONENT_FULL) -> None: + component: radixshmem.ComponentType = COMPONENT_FULL) -> None: self._engine = engine self._sequence_meta = sequence_meta self._slots = slots @@ -158,7 +158,7 @@ def __init__(self, metrics_collector=None): """`server_name` is the radix-server's ``--name``; the server is the operator's process, running but not necessarily ready. With - `geometry` (FlexKV's `RadixGeometry` or a `shmradix.Geometry`) the + `geometry` (FlexKV's `RadixGeometry` or a `radixshmem.Geometry`) the attach hands it FlexKV's slot shape (idempotent) and waits for it to come up. `peer_enabled` None = follow the region: peer reuse whenever the server is part of a cluster; False switches it off. @@ -303,7 +303,7 @@ def insert(self, sequence_meta: SequenceMeta, physical_block_ids: np.ndarray, num_insert_blocks: int, - component: shmradix.ComponentType = COMPONENT_FULL) -> None: + component: radixshmem.ComponentType = COMPONENT_FULL) -> None: """Attach transferred slots: `physical_block_ids[i]` is block `num_insert_blocks - len(physical_block_ids) + i`. Ownership passes to radixshmem (`auto_recycle=True`); do not recycle these slots again.""" @@ -328,12 +328,12 @@ def insert(self, auto_recycle=True, component=component) landed = num_slots - len(result.unused_slots) - if result.error == shmradix.InsertError.FULL_PATH_MISSING: + if result.error == radixshmem.InsertError.FULL_PATH_MISSING: flexkv_logger.warning( f"radixshmem {component} insert on {self.shm_name}: full path " f"[0, {path_end}) was evicted before the window published " f"(slots were auto-recycled)") - elif result.error != shmradix.InsertError.OK: + elif result.error != radixshmem.InsertError.OK: flexkv_logger.warning( f"radixshmem {component} insert on {self.shm_name} returned " f"{result.error}: {landed}/{num_slots} blocks landed at " @@ -352,7 +352,7 @@ def insert(self, f"locally but the cluster-wide publish failed: {e}") if (self.event_collector is not None and component == COMPONENT_FULL - and result.error == shmradix.InsertError.OK): + and result.error == radixshmem.InsertError.OK): # Error-free, the only unused slots are a redundant prefix, so what # landed is the tail of the path. self.event_collector.publish_stored( @@ -362,7 +362,7 @@ def insert(self, def take(self, num_required_blocks: int, - component: shmradix.ComponentType = COMPONENT_FULL) -> np.ndarray: + component: radixshmem.ComponentType = COMPONENT_FULL) -> np.ndarray: """Allocate up to `num_required_blocks` slots, evicting unpinned LRU blocks as needed; fewer come back when the pool cannot supply them (the SWA pool is all-or-none). A request above the pool's size is @@ -390,7 +390,7 @@ def take(self, self._metrics_collector.record_allocation("cpu", len(slots)) return slots - def _pool_total(self, component: shmradix.ComponentType) -> Optional[int]: + def _pool_total(self, component: radixshmem.ComponentType) -> Optional[int]: """Slots in `component`'s pool, None when the index has no such pool.""" if component == COMPONENT_FULL: return int(self._tree.mempool_total()) @@ -400,7 +400,7 @@ def _pool_total(self, component: shmradix.ComponentType) -> Optional[int]: def recycle(self, physical_blocks: np.ndarray, - component: shmradix.ComponentType = COMPONENT_FULL) -> None: + component: radixshmem.ComponentType = COMPONENT_FULL) -> None: if physical_blocks is None or len(physical_blocks) == 0: return self._tree.recycle_slots(np.ascontiguousarray(physical_blocks, dtype=np.int32), diff --git a/flexkv/kvtask.py b/flexkv/kvtask.py index 73e78cab8..6e95d9ce7 100644 --- a/flexkv/kvtask.py +++ b/flexkv/kvtask.py @@ -172,7 +172,7 @@ def __init__(self, f"[KVTaskEngine] topology: {self.model_config}" ) - # radixshmem prefetch jobs in flight: task_id -> shmradix PullJob. Polled + # radixshmem prefetch jobs in flight: task_id -> radixshmem PullJob. Polled # in _update_tasks, the thread every other task mutation runs on. self.prefetch_jobs: Dict[int, Any] = {} if GLOBAL_CONFIG_FROM_ENV.enable_radixshmem: diff --git a/flexkv/server/server.py b/flexkv/server/server.py index 19176f29f..c2f56492c 100644 --- a/flexkv/server/server.py +++ b/flexkv/server/server.py @@ -257,7 +257,7 @@ def create_server(cls, if key.startswith("FLEXKV_") and key not in env: env[key] = val # The child runs the parent's interpreter and must import what - # the parent imports (flexkv, shmradix in radixshmem mode) and + # the parent imports (flexkv, radixshmem in radixshmem mode) and # find the same shared libraries; these are the variables that # locate them when the packages are not installed into # site-packages. diff --git a/flexkv/server/shm_radix_bootstrap.py b/flexkv/server/shm_radix_bootstrap.py index 5211fe600..99ffd5ae5 100644 --- a/flexkv/server/shm_radix_bootstrap.py +++ b/flexkv/server/shm_radix_bootstrap.py @@ -44,9 +44,9 @@ from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType try: - import shmradix + import radixshmem except ImportError: # pragma: no cover - shmradix = None + radixshmem = None # Pool bases are page aligned regardless; a larger per-slot alignment only pads. @@ -106,15 +106,15 @@ def default_endpoint(server_name: str) -> str: return f"unix:///dev/shm/{server_name.lstrip('/').replace('/', '_')}.sock" -def _ensure_shmradix() -> None: - if shmradix is None: +def _ensure_radixshmem() -> None: + if radixshmem is None: raise ImportError( - "shmradix is not installed; install it from the radixshmem repo " + "radixshmem is not installed; install it from the radixshmem repo " "(pip install -e radixshmem/python)") for name in ("RadixClient", "Geometry", "GeometryMismatch", "ServerNotReady"): - if not hasattr(shmradix, name): + if not hasattr(radixshmem, name): raise ImportError( - f"shmradix lacks {name}: FlexKV needs a radixshmem whose radix-server takes " + f"radixshmem lacks {name}: FlexKV needs a radixshmem whose radix-server takes " f"its geometry from the client (RadixClient(name, Geometry))") @@ -231,11 +231,11 @@ def has_swa(self) -> bool: def slot_align(self) -> int: return slot_align_for(self.full_slot_bytes, self.swa_slot_bytes) - def to_shmradix(self) -> "shmradix.Geometry": - """The ``shmradix.Geometry`` handed to the server (data mode: bytes and + def to_radixshmem(self) -> "radixshmem.Geometry": + """The ``radixshmem.Geometry`` handed to the server (data mode: bytes and window only, no counts).""" - _ensure_shmradix() - return shmradix.Geometry( + _ensure_radixshmem() + return radixshmem.Geometry( block_size=int(self.tokens_per_block), full_slot_bytes=int(self.full_slot_bytes), swa_slot_bytes=int(self.swa_slot_bytes), @@ -296,7 +296,7 @@ def _nothing_answers(e: BaseException) -> bool: isinstance(e, RuntimeError) and str(e).startswith("UNAVAILABLE")) -def _current_status(client: "shmradix.RadixClient"): +def _current_status(client: "radixshmem.RadixClient"): """The server's state now (one RPC); the constructor-time info when the server cannot be asked any more.""" try: @@ -311,8 +311,8 @@ def attach_radix_client(name: Optional[str] = None, timeout_s: Optional[float] = None, max_outstanding: Optional[int] = None, attach_index: bool = False, - label: str = "radixshmem") -> "shmradix.RadixClient": - """A ready ``shmradix.RadixClient`` on the radix-server ``name`` (default: + label: str = "radixshmem") -> "radixshmem.RadixClient": + """A ready ``radixshmem.RadixClient`` on the radix-server ``name`` (default: :func:`radix_server_name`), at the socket radixshmem derives from the name. ``attach_index=True`` also brings the index client up before returning @@ -320,7 +320,7 @@ def attach_radix_client(name: Optional[str] = None, peer that does not accept yet; callers that only read ``info`` (adopting counts, asking about the cluster) leave it off and open no RDMA state. - With ``geometry`` (a :class:`RadixGeometry` or a ``shmradix.Geometry``) the + With ``geometry`` (a :class:`RadixGeometry` or a ``radixshmem.Geometry``) the client hands the server FlexKV's slot shape on the way; the server plans the counts from its budget. That is idempotent, so every FlexKV process may bring it; a server already serving another geometry (another model or @@ -333,22 +333,22 @@ def attach_radix_client(name: Optional[str] = None, the SlotStore prefault happen there -- for ``timeout_s`` in total (default ``READY_TIMEOUT_S``). """ - _ensure_shmradix() + _ensure_radixshmem() name = name or radix_server_name() if timeout_s is None: timeout_s = READY_TIMEOUT_S if max_outstanding is None: max_outstanding = MAX_OUTSTANDING - spec = geometry.to_shmradix() if isinstance(geometry, RadixGeometry) else geometry + spec = geometry.to_radixshmem() if isinstance(geometry, RadixGeometry) else geometry where = default_endpoint(name) deadline = time.monotonic() + float(timeout_s) last: Optional[BaseException] = None while True: try: - client = shmradix.RadixClient(name, spec, max_outstanding=max_outstanding) + client = radixshmem.RadixClient(name, spec, max_outstanding=max_outstanding) break - except shmradix.GeometryMismatch as e: + except radixshmem.GeometryMismatch as e: raise ValueError( f"{label}: radix-server {name} already serves another geometry ({e}); every " f"engine attached to one server must run the same model, page size and SWA " @@ -397,7 +397,7 @@ def attach_radix_client(name: Optional[str] = None, return client -def _attach_index(client: "shmradix.RadixClient", name: str, deadline: float, label: str) -> None: +def _attach_index(client: "radixshmem.RadixClient", name: str, deadline: float, label: str) -> None: """Bring the index client up now (``client.index``) instead of on first use. On a cluster this opens the RDMA queue pairs to every peer's RHT shard; right after the rendezvous a peer's holder may not accept yet, and @@ -434,7 +434,7 @@ def _describe_published(g: Optional[Dict[str, Any]]) -> str: return ", ".join(parts) -def _published_geometry(client: "shmradix.RadixClient", label: str) -> Dict[str, Any]: +def _published_geometry(client: "radixshmem.RadixClient", label: str) -> Dict[str, Any]: g = client.geometry if not g: info = client.status() @@ -446,13 +446,13 @@ def _published_geometry(client: "shmradix.RadixClient", label: str) -> Dict[str, return g -def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, +def check_geometry(client: "radixshmem.RadixClient", expected: RadixGeometry, label: str = "radixshmem") -> None: """Fail closed when the server's regions differ from FlexKV's own layout: a stride or page mismatch would otherwise become a silent misaddressed transfer. Slot counts are not checked here -- they are the server's, taken over by :func:`adopt_geometry`.""" - _ensure_shmradix() + _ensure_radixshmem() g = _published_geometry(client, label) pools = g["pools"] diffs: List[str] = [] @@ -474,7 +474,7 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, diffs.append("server is index-only (no --data-bytes); FlexKV needs the data plane") else: store = client.store - stride = int(store.pool(shmradix.ComponentType.FULL).slot_bytes) + stride = int(store.pool(radixshmem.ComponentType.FULL).slot_bytes) if stride != expected.full_slot_bytes: diffs.append(f"FULL stride server={stride} flexkv={expected.full_slot_bytes} " f"(the server rounds slots up to slot_align={g.get('slot_align')}; " @@ -491,7 +491,7 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, diffs.append(f"SWA window server={swa.get('window_blocks')} " f"flexkv={expected.swa_window_blocks}") if client.info.data_plane: - stride = int(client.store.pool(shmradix.ComponentType.SWA).slot_bytes) + stride = int(client.store.pool(radixshmem.ComponentType.SWA).slot_bytes) if stride != expected.swa_slot_bytes: diffs.append(f"SWA stride server={stride} flexkv={expected.swa_slot_bytes}") elif swa is not None: @@ -502,7 +502,7 @@ def check_geometry(client: "shmradix.RadixClient", expected: RadixGeometry, f"({expected.describe()}): " + "; ".join(diffs)) -def adopt_geometry(cache_config: CacheConfig, client: "shmradix.RadixClient", +def adopt_geometry(cache_config: CacheConfig, client: "radixshmem.RadixClient", label: str = "radixshmem") -> Dict[str, int]: """Take over what the server planned and published: the slot counts into ``cache_config`` (``num_cpu_blocks`` = the FULL pool, ``swa.num_slots`` = @@ -585,7 +585,7 @@ def radix_server_is_distributed(model_config: ModelConfig, cache_config: CacheCo client.close() -def radix_cluster_rank(client: "shmradix.RadixClient") -> int: +def radix_cluster_rank(client: "radixshmem.RadixClient") -> int: """This node's rank in the radix cluster (0 on a standalone server).""" rank = int(getattr(client.info, "rank", -1)) return rank if rank >= 0 else int(client.rank()) diff --git a/flexkv/storage/allocator.py b/flexkv/storage/allocator.py index 0f7ea35b6..b4cad1e93 100644 --- a/flexkv/storage/allocator.py +++ b/flexkv/storage/allocator.py @@ -327,12 +327,12 @@ class SlotStoreTensorHandle: """ data_name: str hugepage_path: str - kind: int # shmradix.ComponentType value (0 = FULL, 1 = SWA) + kind: int # radixshmem.ComponentType value (0 = FULL, 1 = SWA) num_elements: int dtype: torch.dtype def get_tensor(self) -> torch.Tensor: - from shmradix import _data + from radixshmem import _data store = _data.SlotStore.attach(self.data_name, self.hugepage_path, 60000) return slot_store_pool_tensor(store, self.kind, self.dtype, self.num_elements) diff --git a/flexkv/storage/storage_engine.py b/flexkv/storage/storage_engine.py index 4e82a7f35..d288574c8 100644 --- a/flexkv/storage/storage_engine.py +++ b/flexkv/storage/storage_engine.py @@ -84,7 +84,7 @@ def __init__(self, radix_client: Any = None): """Initialize storage engine. - ``radix_client`` (a ``shmradix.RadixClient``, radixshmem mode) makes the + ``radix_client`` (a ``radixshmem.RadixClient``, radixshmem mode) makes the CPU FULL / SWA pools views of the radix-server's SlotStore instead of allocations of this process; see ``_attach_radix_pool``. """ @@ -326,7 +326,7 @@ def _attach_radix_pool(self, misaddressed transfer, not an error. Workers re-attach the pool by name through the ``SlotStoreTensorHandle`` in ``worker_data``. """ - from shmradix import ComponentType + from radixshmem import ComponentType from flexkv.server.shm_radix_bootstrap import layout_block_bytes kind = ComponentType.SWA if is_swa else ComponentType.FULL diff --git a/tests/radixshmem/radix_e2e_common.py b/tests/radixshmem/radix_e2e_common.py index d3c8da81a..485d2d164 100644 --- a/tests/radixshmem/radix_e2e_common.py +++ b/tests/radixshmem/radix_e2e_common.py @@ -94,11 +94,11 @@ def stop_private_etcd(proc, workdir) -> None: def start_radix_server(name: str, data_bytes: int, *, extra_args=(), endpoint: Optional[str] = None, log_path: Optional[str] = None, timeout: float = 60.0) -> subprocess.Popen: - """Start the operator's ``radix-server`` (``python -m shmradix.cli``) and wait + """Start the operator's ``radix-server`` (``python -m radixshmem.cli``) and wait for its socket. Nothing model-specific goes on its command line: FlexKV's clients bring the geometry, the server plans the slot counts from ``data_bytes``.""" - cmd = [sys.executable, "-m", "shmradix.cli", "--name", name, "--data-bytes", str(int(data_bytes)), + cmd = [sys.executable, "-m", "radixshmem.cli", "--name", name, "--data-bytes", str(int(data_bytes)), "--no-prefault", "--interval", "0", *extra_args] if endpoint: cmd += ["--endpoint", endpoint] diff --git a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py index 454b8b76f..fe491c9f6 100644 --- a/tests/radixshmem/test_e2e_radix_prefetch_p2p.py +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -22,7 +22,7 @@ LOCAL_HEAD_BLOCKS of it (its own bytes), the prefetch pulls only the tail, and the GET serves head and tail from the right writers. -Requires >=2 CUDA devices, an ACTIVE RDMA port, a shmradix built with RDMA + +Requires >=2 CUDA devices, an ACTIVE RDMA port, a radixshmem built with RDMA + etcd + mooncake, and an etcd (FLEXKV_TEST_RADIX_REGISTRY, or ``etcd`` on PATH for a private one); skips otherwise. Run inside the container: @@ -290,7 +290,7 @@ def _run(registry: str, rdma_dev: str) -> dict: def cluster(): """(etcd registry, rdma device), skipping when the prerequisites are absent; starts a private etcd when none is configured.""" - pytest.importorskip("shmradix") + pytest.importorskip("radixshmem") if not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE: pytest.skip(f"needs {WORLD_SIZE} CUDA devices") devices = active_rdma_devices() diff --git a/tests/radixshmem/test_e2e_radix_shmem.py b/tests/radixshmem/test_e2e_radix_shmem.py index 2fb237508..d82e2e417 100644 --- a/tests/radixshmem/test_e2e_radix_shmem.py +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -235,7 +235,7 @@ def _run(dp_size: int) -> dict: @pytest.mark.e2e @pytest.mark.parametrize("dp_size", [1, 2]) def test_radix_shmem_put_get_roundtrip(dp_size): - pytest.importorskip("shmradix") + pytest.importorskip("radixshmem") if not torch.cuda.is_available() or torch.cuda.device_count() < dp_size: pytest.skip(f"needs {dp_size} CUDA device(s)") diff --git a/tests/radixshmem/test_radix_shmem_engine.py b/tests/radixshmem/test_radix_shmem_engine.py index b35c6f12e..98a824d89 100644 --- a/tests/radixshmem/test_radix_shmem_engine.py +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -1,11 +1,11 @@ """Tests for the radixshmem CPU tier: the engine on a radix-server, the data plane FlexKV maps as its CPU pool, the planners, and the peer pull. -Skipped if `shmradix` is missing. +Skipped if `radixshmem` is missing. Four parts, one file: Part 1 — `CacheEngineRadixShmem` semantics against an in-process - `shmradix.RadixServer` (index + SlotStore, started the way the operator's + `radixshmem.RadixServer` (index + SlotStore, started the way the operator's `radix-server` is: a name and a byte budget; FlexKV's client brings the geometry): take / insert / match / recycle, insert-after-transfer publication, lock vs. eviction, SWA windows, and the fact that a standalone @@ -27,7 +27,7 @@ `radix_shmem_engine.py`, so they do not pull in `flexkv.c_ext` (CUDA). * Part 2 needs the real `GlobalCacheEngine` (and therefore `c_ext`), so it imports it lazily and skips instead of breaking collection for Parts 1/3. - * Part 3 needs an ACTIVE RDMA device, a shmradix built WITH RDMA + etcd + + * Part 3 needs an ACTIVE RDMA device, a radixshmem built WITH RDMA + etcd + mooncake, and an etcd (FLEXKV_TEST_RADIX_REGISTRY, or an `etcd` binary on PATH to start a private one); it is gated behind FLEXKV_RUN_RADIX_PEER_TEST=1. """ @@ -55,18 +55,18 @@ import pytest try: - import shmradix + import radixshmem except ImportError as exc: - # Not importorskip: a shmradix built before the current API (stale `_core.so` + # Not importorskip: a radixshmem built before the current API (stale `_core.so` # next to a newer `__init__.py`) raises ImportError rather than # ModuleNotFoundError, and pytest >= 8.2 only skips on the latter — which # would abort collection for the whole suite instead of skipping this file. - pytest.skip(f"shmradix unusable ({exc}); rebuild the extension", + pytest.skip(f"radixshmem unusable ({exc}); rebuild the extension", allow_module_level=True) for _name in ("RadixServer", "ServerConfig", "Geometry", "RadixClient"): - if not hasattr(shmradix, _name): - pytest.skip(f"shmradix lacks {_name}: needs the RadixServer/RadixClient surface", + if not hasattr(radixshmem, _name): + pytest.skip(f"radixshmem lacks {_name}: needs the RadixServer/RadixClient surface", allow_module_level=True) @@ -106,8 +106,8 @@ def _load_module_direct(name: str, path: str): from flexkv.common.transfer import TransferType # noqa: E402 from flexkv.server import shm_radix_bootstrap as bootstrap # noqa: E402 -FULL = shmradix.ComponentType.FULL -_SWA = shmradix.ComponentType.SWA +FULL = radixshmem.ComponentType.FULL +_SWA = radixshmem.ComponentType.SWA SLOT_BYTES = 256 # bytes of one test block in the SlotStore @@ -163,9 +163,9 @@ def _server_config(name: str, blocks: int, tokens_per_block: int, server plans equal the test's ``blocks`` / ``swa_slots``.""" align = bootstrap.slot_align_for(slot_bytes) data_bytes, swa_ratio = _server_budget(blocks, swa_slots, slot_bytes) - cfg = shmradix.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=swa_ratio, + cfg = radixshmem.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=swa_ratio, slot_align=align, prefault=False) - geo = shmradix.Geometry(block_size=tokens_per_block, full_slot_bytes=slot_bytes, + geo = radixshmem.Geometry(block_size=tokens_per_block, full_slot_bytes=slot_bytes, swa_slot_bytes=slot_bytes if swa_slots else 0, swa_window_blocks=window_blocks if swa_slots else 0, slot_align=align) @@ -182,7 +182,7 @@ def __init__(self) -> None: def server(self, cfg): """Start the server: ``waiting`` until a client brings the geometry.""" _sweep_region(cfg.name, cfg.resolved_data_name) - server = shmradix.RadixServer(cfg).start() + server = radixshmem.RadixServer(cfg).start() self._stack.append(server.close) return server @@ -653,13 +653,13 @@ def test_expected_geometry_mirrors_the_storage_engine_layout(): assert geo.tokens_per_block == 16 and geo.full_slot_bytes == 32768 assert geo.has_swa and geo.swa_slot_bytes == 1024 and geo.swa_window_blocks == SWA_W assert geo.slot_align == 1024 # gcd power of two of 32768 and 1024 - spec = geo.to_shmradix().to_dict() + spec = geo.to_radixshmem().to_dict() assert spec["block_size"] == 16 and spec["slot_align"] == 1024 assert spec["pools"]["full"] == {"slot_bytes": 32768, "num_slots": 0} assert spec["pools"]["swa"] == {"slot_bytes": 1024, "num_slots": 0, "window_blocks": SWA_W} # Without SWA the pool is absent from what the server is asked for. model_config, cache_config = _configs(num_cpu_blocks=64) - spec = bootstrap.expected_geometry(model_config, cache_config).to_shmradix().to_dict() + spec = bootstrap.expected_geometry(model_config, cache_config).to_radixshmem().to_dict() assert set(spec["pools"]) == {"full"} @@ -671,10 +671,10 @@ def test_register_chunk_is_the_servers_unless_pinned(): model_config, cache_config = _configs(num_cpu_blocks=64) geo = bootstrap.expected_geometry(model_config, cache_config) assert geo.register_chunk_tokens == 0 - assert geo.to_shmradix().to_dict()["register_chunk_tokens"] == 0 + assert geo.to_radixshmem().to_dict()["register_chunk_tokens"] == 0 assert "register_chunk" not in geo.describe() pinned = dataclasses.replace(geo, register_chunk_tokens=2048) - assert pinned.to_shmradix().to_dict()["register_chunk_tokens"] == 2048 + assert pinned.to_radixshmem().to_dict()["register_chunk_tokens"] == 2048 assert "register_chunk_tokens=2048" in pinned.describe() assert bootstrap.register_chunk_blocks(4096, 16) == 256 assert bootstrap.register_chunk_blocks(4096, 4) == 1024 @@ -691,7 +691,7 @@ def test_client_brings_the_geometry_and_adopts_the_counts(env): geo = bootstrap.expected_geometry(model_config, cache_config) name = f"/geo{os.getpid()}" data_bytes = 64 * geo.full_slot_bytes + 16 * geo.swa_slot_bytes - env.server(shmradix.ServerConfig(name=name, data_bytes=data_bytes, + env.server(radixshmem.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=16 * geo.swa_slot_bytes / data_bytes, register_chunk_tokens=2048, # the operator's, not FlexKV's prefault=False)) # slot_align comes with the geometry @@ -742,7 +742,7 @@ def test_attach_waits_for_a_late_server(env): def _start_later(): time.sleep(1.5) - env._stack.append(shmradix.RadixServer(cfg).start().close) + env._stack.append(radixshmem.RadixServer(cfg).start().close) started.set() threading.Thread(target=_start_later, daemon=True).start() @@ -776,13 +776,13 @@ def _ctor(*args, **kwargs): (TypeError("unexpected keyword argument 'endpoint'"), TypeError, "unexpected keyword"), (RuntimeError("INTERNAL: server is closed"), RuntimeError, "server is closed")): calls.clear() - monkeypatch.setattr(shmradix, "RadixClient", _ctor_raising(exc)) + monkeypatch.setattr(radixshmem, "RadixClient", _ctor_raising(exc)) t0 = time.monotonic() with pytest.raises(exc_type, match=match): bootstrap.attach_radix_client(name, timeout_s=30) assert len(calls) == 1 and time.monotonic() - t0 < 5 calls.clear() - monkeypatch.setattr(shmradix, "RadixClient", + monkeypatch.setattr(radixshmem, "RadixClient", _ctor_raising(RuntimeError("UNAVAILABLE: failed to connect to all addresses"))) with pytest.raises(TimeoutError, match="no radix-server named"): bootstrap.attach_radix_client(name, timeout_s=1.5) @@ -821,16 +821,16 @@ def _ctor(name, spec, max_outstanding=256): return fakes[-1] monkeypatch.setattr(bootstrap.time, "sleep", lambda s: None) - monkeypatch.setattr(shmradix, "RadixClient", _ctor) + monkeypatch.setattr(radixshmem, "RadixClient", _ctor) client = bootstrap.attach_radix_client("/x", timeout_s=60, attach_index=True) assert client is fakes[0] and fakes[0].failures == 0 and not fakes[0].closed fakes.clear() - monkeypatch.setattr(shmradix, "RadixClient", lambda *a, **k: _Fake(1, "index/store/geometry mismatch: x")) + monkeypatch.setattr(radixshmem, "RadixClient", lambda *a, **k: _Fake(1, "index/store/geometry mismatch: x")) with pytest.raises(RuntimeError, match="attaching the index failed"): bootstrap.attach_radix_client("/x", timeout_s=60, attach_index=True) # without attach_index the index is left alone (info-only callers open no RDMA state) fakes.clear() - monkeypatch.setattr(shmradix, "RadixClient", lambda *a, **k: _Fake(5)) + monkeypatch.setattr(radixshmem, "RadixClient", lambda *a, **k: _Fake(5)) assert bootstrap.attach_radix_client("/x", timeout_s=60).failures == 5 @@ -841,7 +841,7 @@ def test_unconfigured_server_answers_the_cluster_question(env): model_config, cache_config = _configs(num_cpu_blocks=64) geo = bootstrap.expected_geometry(model_config, cache_config) name = f"/waiting{os.getpid()}" - env.server(shmradix.ServerConfig(name=name, data_bytes=64 * geo.full_slot_bytes, + env.server(radixshmem.ServerConfig(name=name, data_bytes=64 * geo.full_slot_bytes, prefault=False)) with pytest.raises(TimeoutError, match="mode=waiting; nobody handed it a geometry"): bootstrap.attach_radix_client(name, timeout_s=2) @@ -908,7 +908,7 @@ def test_radix_server_name_rejects_malformed_names(monkeypatch, name): class FakeJob: - """Stand-in for `shmradix.PullJob`: what `_plan_prefetch` reads on return + """Stand-in for `radixshmem.PullJob`: what `_plan_prefetch` reads on return (`local_hit`, `planned_hit`) and what `KVTaskEngine` polls.""" def __init__(self, local_hit: int, planned_hit: int, job_id: int = 7): @@ -1465,11 +1465,11 @@ def _swa_global_engine(swa_slots: int = 2 * SWA_W, # test's; the planner's client brings the geometry. geo = bootstrap.expected_geometry(model_config, cache_config) data_bytes = num_blocks * geo.full_slot_bytes + swa_slots * geo.swa_slot_bytes - cfg = shmradix.ServerConfig(name=server_name, data_bytes=data_bytes, + cfg = radixshmem.ServerConfig(name=server_name, data_bytes=data_bytes, swa_ratio=swa_slots * geo.swa_slot_bytes / data_bytes, prefault=False) _sweep_region(cfg.name, cfg.resolved_data_name) - server = shmradix.RadixServer(cfg).start() + server = radixshmem.RadixServer(cfg).start() engine = RadixShmemCacheEngine(cache_config, model_config) assert engine.swa_op_constructor.enabled, \ "SWA gate should be on: enable_swa_transfer + radixshmem swa_enabled" @@ -1535,7 +1535,7 @@ def _recording_insert(*args, **kwargs): pending.release() put_cb() # graph completion - assert published == [shmradix.ComponentType.FULL, _SWA] + assert published == [radixshmem.ComponentType.FULL, _SWA] after = cpu.match(_real_seq(token_ids), component_mask=JOINT_MASK) assert after.num_matched_blocks == 20 @@ -1727,7 +1727,7 @@ def test_planner_uses_configured_window_blocks(): # Two spawned processes each run a data-mode RadixServer (distinct data names # and sockets on one host) and a `CacheEngineRadixShmem` attached to it. # Gated behind FLEXKV_RUN_RADIX_PEER_TEST=1; needs an ACTIVE RDMA device, a -# shmradix built with RDMA + etcd + mooncake, and an etcd +# radixshmem built with RDMA + etcd + mooncake, and an etcd # (FLEXKV_TEST_RADIX_REGISTRY, or `etcd` on PATH for a private one). # ============================================================================= @@ -1767,11 +1767,11 @@ def cluster(): if not devices: pytest.skip("no ACTIVE RDMA device found") try: - from shmradix import _data + from radixshmem import _data if not hasattr(_data, "DataPlaneRegistry"): - pytest.skip("shmradix built without etcd (no DataPlaneRegistry)") + pytest.skip("radixshmem built without etcd (no DataPlaneRegistry)") except ImportError: - pytest.skip("shmradix built without the _data extension") + pytest.skip("radixshmem built without the _data extension") registry = os.getenv("FLEXKV_TEST_RADIX_REGISTRY", "") proc = None workdir = None @@ -1828,15 +1828,15 @@ def _node_main(rank, prefix, cluster_id, registry, rdma_dev, ready, done, output node_name=f"r{rank}", rpc_address="0.0.0.0", index_dev=rdma_dev, gid_idx=int(os.getenv("FLEXKV_TEST_RADIX_GID_IDX", "3")), bootstrap_timeout_sec=60, rht_slots_per_bucket=4) - cfg = shmradix.ServerConfig( + cfg = radixshmem.ServerConfig( name=name, data_bytes=PEER_BLOCKS * PEER_SLOT_BYTES, slot_align=4096, data_name=data_name, prefault=False, transfer_devices=[rdma_dev], - cluster=shmradix.ClusterConfig(**cluster_kwargs), + cluster=radixshmem.ClusterConfig(**cluster_kwargs), ) - server = shmradix.RadixServer(cfg).start() # waiting: the engine's geometry starts the rendezvous + server = radixshmem.RadixServer(cfg).start() # waiting: the engine's geometry starts the rendezvous bootstrap.READY_TIMEOUT_S = 180.0 # this process only: the rendezvous may take a while engine = CacheEngineRadixShmem( - name, geometry=shmradix.Geometry(block_size=16, full_slot_bytes=PEER_SLOT_BYTES, + name, geometry=radixshmem.Geometry(block_size=16, full_slot_bytes=PEER_SLOT_BYTES, slot_align=4096), num_total_blocks=PEER_BLOCKS, tokens_per_block=16, peer_enabled=True) if not engine.peer_enabled: @@ -1919,7 +1919,7 @@ def _run_two_nodes(registry, rdma_dev, local_head_blocks=0): ready = ctx.Event() done = ctx.Event() output = ctx.Queue() - prefix = f"/shmradix_peer_test_{os.getpid()}_{local_head_blocks}" + prefix = f"/radixshmem_peer_test_{os.getpid()}_{local_head_blocks}" cluster_id = f"flexkv-peer-test-{os.getpid()}-{local_head_blocks}" processes = [