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..32865c913 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,11 @@ 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](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. + 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/csrc/bindings.cpp b/csrc/bindings.cpp index d7a33b19a..e254eb7d5 100644 --- a/csrc/bindings.cpp +++ b/csrc/bindings.cpp @@ -10,6 +10,7 @@ #include #include #include +#include #include #include #include @@ -481,6 +482,48 @@ 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. + // 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++) { + 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 +809,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/docs/radixshmem/config_zh.md b/docs/radixshmem/config_zh.md new file mode 100644 index 000000000..e7025410d --- /dev/null +++ b/docs/radixshmem/config_zh.md @@ -0,0 +1,137 @@ +# radixshmem 模式配置 + +FlexKV 以 radixshmem 作为 CPU 层(索引 + SlotStore + 跨节点拉取)。配置分两处: + +| 谁 | 载体 | 内容 | +|---|---|---| +| 运维 | `radix-server` 命令行,每节点一个进程 | 名字、SlotStore 字节预算与 SWA 占比、hugepage、传输引擎、集群成员(etcd、网卡、rank)、索引调优 | +| FlexKV | 两个环境变量 | 是否启用、attach 哪个 server | + +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 提示)。 + +实现:`flexkv/server/shm_radix_bootstrap.py`(几何、attach、采纳、固定参数)。 +radixshmem 侧接口见 radixshmem 仓库 `python/README.md`。 + +## 1. 环境变量 + +| 变量 | 默认 | 说明 | +|---|---|---| +| `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 节)。 | + +该模式只承担 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. 固定参数 + +attach 的其余参数是 `flexkv/server/shm_radix_bootstrap.py` 里的常量: + +| 常量 | 值 | 含义 | +|---|---|---| +| `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 数 + +FlexKV 交给 server 的几何(`shm_radix_bootstrap.expected_geometry` → `radixshmem.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 数 | + +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`。 + +**采纳**(`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`。 + +**校验**(`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.1 单机 + +```bash +# 运维,每节点一次;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_SERVER_NAME 即 attach /flexkv +``` + +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`: + +```bash +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_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`)由第一个 +拿到几何的节点发布到 etcd `radix//geometry/`,其余节点采纳;各节点的 slot 数可以不同。FlexKV 侧每个节点 +同一个 `FLEXKV_RADIXSHMEM_SERVER_NAME`。 + +### 4.3 同机多节点(测试) + +两个 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: + +```bash +# 引擎 A # 引擎 B +FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=0 FLEXKV_INSTANCE_NUM=2 FLEXKV_INSTANCE_ID=1 +``` + +同一个 `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. 命名 + +| 对象 | 名字 | +|---|---| +| index shm | `--name`;集群模式下 radixshmem 追加 `_`,attach 方只需 `--name` | +| SlotStore shm | `_data`(`--data-name` 可改) | +| gRPC socket | `/dev/shm/.sock`,FlexKV 按名字派生;server 端保持默认 `--endpoint` | +| etcd 键空间 | `radix//...` | + +## 6. 报错含义 + +- `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 有响应但 attach + 失败(如 `server is closed`)立即报错;server 一直在等几何或配置失败报 `not ready within ...`(附 server 当时的 + `mode` 和 `last_error`);几何冲突见第 3 节。 diff --git a/examples/radixshmem/radix_server_multi_node.sh b/examples/radixshmem/radix_server_multi_node.sh new file mode 100755 index 000000000..a0894e881 --- /dev/null +++ b/examples/radixshmem/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/radix_server_single_node.sh b/examples/radixshmem/radix_server_single_node.sh new file mode 100755 index 000000000..4fe94bf5e --- /dev/null +++ b/examples/radixshmem/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/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..fafff44a3 --- /dev/null +++ b/flexkv/cache/radix_shmem_engine.py @@ -0,0 +1,407 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +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`. + +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 radixshmem +except ImportError as e: # pragma: no cover + raise ImportError( + "radixshmem is not installed; install it from the radixshmem repo " + "(pip install -e radixshmem/python)") from e + +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: + 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: radixshmem.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, + server_name: str, + *, + tokens_per_block: int, + num_total_blocks: int, + geometry: Any = None, + peer_enabled: Optional[bool] = None, + swa_config: Optional[SWAPoolConfig] = None, + event_collector: Optional[KVEventCollector] = None, + 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 `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. + `num_total_blocks` is FlexKV's expectation; the region's capacity is + authoritative.""" + from flexkv.server.shm_radix_bootstrap import attach_radix_client, register_chunk_blocks + + 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(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 + 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( + 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) + # 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( + 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: 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.""" + 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 == 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 != radixshmem.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: + # 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 == 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( + 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: 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 + 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: 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()) + if component == COMPONENT_SWA: + return int(self._tree.swa_mempool_total()) + return None + + def recycle(self, + physical_blocks: np.ndarray, + 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), + component=component) diff --git a/flexkv/cache/radix_shmem_planner.py b/flexkv/cache/radix_shmem_planner.py new file mode 100644 index 000000000..8c65eb83b --- /dev/null +++ b/flexkv/cache/radix_shmem_planner.py @@ -0,0 +1,602 @@ +# 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 radix-server instead (on whenever it is part of a cluster). +""" + +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.server.shm_radix_bootstrap import (PREFETCH_MAX_INFLIGHT, PREFETCH_TIMEOUT_MS, + expected_geometry, radix_server_name) +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 it is started with cluster flags); " + "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) + # GetJobs this engine started and has not yet seen finish; pruned on + # 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) + + # ------------------------------------------------------------------ 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 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. + """ + return CacheEngineRadixShmem( + 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, + 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() + inflight = self._prefetch_inflight() + if inflight >= PREFETCH_MAX_INFLIGHT: + flexkv_logger.debug( + f"radixshmem prefetch {request_id}: {inflight} peer pulls in flight " + 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 + job = engine.prefetch( + sequence_meta, + component_mask=mask, + query_end=block_mask_end, + timeout_ms=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 + 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, + ) + 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/config.py b/flexkv/common/config.py index 5573fef2a..e086416f5 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,14 @@ 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 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_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()), ssd_layout_type=KVCacheLayoutType(os.getenv('FLEXKV_SSD_LAYOUT', 'BLOCKFIRST').upper()), @@ -1174,6 +1228,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/hash_utils.py b/flexkv/common/hash_utils.py index a625bb1f6..784205061 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()) @@ -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(hasher.hasher, torch.from_numpy(token_ids), tokens_per_block, torch.from_numpy(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/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( 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/integration/sglang/connector.py b/flexkv/integration/sglang/connector.py index f76508d8a..40da36e4e 100644 --- a/flexkv/integration/sglang/connector.py +++ b/flexkv/integration/sglang/connector.py @@ -61,6 +61,18 @@ from flexkv.transfer.layer_eventfd import build_layerwise_eventfd_socket_path from flexkv.transfer_manager import TransferManagerOnRemote + +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 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(model_config, cache_config, label="FlexKVConnector") + + logger = logging.getLogger(__name__) _SGLANG_REQ_ID_UNSET = object() @@ -167,7 +179,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 +241,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 +276,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 +341,10 @@ 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.model_config, self.cache_config)) ) self._shutdown_done = False @@ -1399,7 +1423,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 +2366,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 +2450,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 +2547,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/flexkv/kvmanager.py b/flexkv/kvmanager.py index fdcebd7f5..76a496def 100644 --- a/flexkv/kvmanager.py +++ b/flexkv/kvmanager.py @@ -38,7 +38,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,11 +61,42 @@ 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 + # 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 " + "radix-server's cluster flags)" + ) 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: + # 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 model_config.instance_num > 1 or @@ -80,18 +112,44 @@ 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" + self.kv_task_engine = None + self.server_handle = None + + if self.enable_radixshmem: + # 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: - if self.server_launch_mode == "embedded" and dp_client_id == 0: + # 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, @@ -161,7 +219,8 @@ 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() if self.owns_mps: flexkv_logger.info( @@ -192,6 +251,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 +285,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 +311,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 +332,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 +362,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..6e95d9ce7 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,7 @@ 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, ): if not cache_config.enable_cpu: raise ValueError("enable_cpu must be True") @@ -167,7 +172,23 @@ 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 -> 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: + # The CPU tier is a radix-server (shared index + SlotStore); its + # 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) if not self.model_config.use_trtllm_subprocess: self.transfer_handles = [TransferManagerHandle( @@ -198,7 +219,10 @@ def __init__(self, ] self.transfer_handles[0]._handle.send_config_to_remotes() - if self.model_config.nnodes > 1: + # 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 # index: release builds cythonize this module with # wraparound=False, so ``transfer_handles[-1]`` on a list reads off @@ -423,6 +447,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 +484,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 +780,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 +923,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,7 +1104,7 @@ 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, ): super().__init__(model_config, cache_config, gpu_register_port, redis_meta, event_collector) self.tracer = FlexKVTracer() @@ -1492,6 +1588,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 diff --git a/flexkv/server/server.py b/flexkv/server/server.py index 438f14740..c2f56492c 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, 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. + 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 new file mode 100644 index 000000000..99ffd5ae5 --- /dev/null +++ b/flexkv/server/shm_radix_bootstrap.py @@ -0,0 +1,591 @@ +# SPDX-License-Identifier: Apache-2.0 +# cython: boundscheck=True, wraparound=True +""" +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 time +from typing import Any, Dict, List, Optional, Tuple + +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.storage import KVCacheLayout, KVCacheLayoutType + +try: + import radixshmem +except ImportError: # pragma: no cover + radixshmem = None + + +# 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 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``.""" + return f"unix:///dev/shm/{server_name.lstrip('/').replace('/', '_')}.sock" + + +def _ensure_radixshmem() -> None: + if radixshmem is None: + raise ImportError( + "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(radixshmem, name): + raise ImportError( + f"radixshmem lacks {name}: FlexKV needs a radixshmem whose radix-server takes " + f"its geometry from the client (RadixClient(name, Geometry))") + + +# ------------------------------------------------------------------ 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]: + """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: + 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: + """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. 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: + return self.swa_slot_bytes > 0 + + @property + def slot_align(self) -> int: + return slot_align_for(self.full_slot_bytes, self.swa_slot_bytes) + + def to_radixshmem(self) -> "radixshmem.Geometry": + """The ``radixshmem.Geometry`` handed to the server (data mode: bytes and + window only, no counts).""" + _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), + 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})" + 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: + 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") + geo = RadixGeometry( + 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) + if swa is not None: + if swa.window_blocks < 1: + raise ValueError( + f"cache_config.swa.window_blocks={swa.window_blocks} must be >= 1") + geo = dataclasses.replace( + geo, swa_slot_bytes=swa_block_bytes(model_config, cache_config), + swa_window_blocks=int(swa.window_blocks)) + return geo + + +# -------------------------------------------------------------------- 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: "radixshmem.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, + timeout_s: Optional[float] = None, + max_outstanding: Optional[int] = None, + attach_index: bool = False, + 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 + (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 ``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 + 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 ``READY_TIMEOUT_S``). + """ + _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_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 = radixshmem.RadixClient(name, spec, max_outstanding=max_outstanding) + break + 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 " + 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 - 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( + 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: + 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}" + + (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 + 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}, " + f"geometry={_describe_published(info.geometry)}") + return client + + +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 + 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)" + 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": + s += f" (window {pool.get('window_blocks')})" + parts.append(s) + return ", ".join(parts) + + +def _published_geometry(client: "radixshmem.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: "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_radixshmem() + 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}") + 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}") + if not client.info.data_plane: + diffs.append("server is index-only (no --data-bytes); FlexKV needs the data plane") + else: + store = client.store + 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')}; " + f"FlexKV asks for {expected.slot_align})") + swa = pools.get("swa") + if expected.has_swa: + if swa is None: + diffs.append("server has no SWA pool but FlexKV's SWA tier is on " + "(start it with --swa-ratio)") + else: + 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(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: + diffs.append("server has an SWA pool that FlexKV's configuration does not") + if diffs: + raise ValueError( + f"{label}: the radix-server's regions do not match FlexKV's layout " + f"({expected.describe()}): " + "; ".join(diffs)) + + +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`` = + 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"])} + 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})" + 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 + + +def adopt_radix_server(model_config: ModelConfig, cache_config: CacheConfig, + *, + 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 + (``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(name, 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(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, 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: + client.close() + + +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 7accf7e8d..b4cad1e93 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 # radixshmem.ComponentType value (0 = FULL, 1 = SWA) + num_elements: int + dtype: torch.dtype + + def get_tensor(self) -> torch.Tensor: + 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) + + +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..d288574c8 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 ``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``. + """ + 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 radixshmem 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, 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..d2d6dd7ab 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,27 @@ 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: + # 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, 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 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 +596,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): """ diff --git a/tests/radixshmem/radix_e2e_common.py b/tests/radixshmem/radix_e2e_common.py new file mode 100644 index 000000000..485d2d164 --- /dev/null +++ b/tests/radixshmem/radix_e2e_common.py @@ -0,0 +1,301 @@ +"""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 sys +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 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 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", "radixshmem.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.""" + 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..fe491c9f6 --- /dev/null +++ b/tests/radixshmem/test_e2e_radix_prefetch_p2p.py @@ -0,0 +1,341 @@ +"""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: + + * 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 + (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 + 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 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: + + 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, + start_radix_server, + stop_radix_server, + sweep_radix_files, + wait_kv_manager_ready, + write_pattern, +) + +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, 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 + # 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. + recv_port = f"ipc:///tmp/flexkv_{cluster_id}_{node_name}" + os.environ.update({ + "FLEXKV_ENABLE_RADIXSHMEM": "1", + "FLEXKV_RADIXSHMEM_SERVER_NAME": server_name, # this node's server + "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_server_name = server_name + 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 radix-server (started with cluster flags + # 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_") + # 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, names = [], [] + for rank in range(WORLD_SIZE): + name = f"/{cluster_id}_{_node_name(rank)}" + names.append(name) + servers.append(start_radix_server( + 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"))) + 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, names[rank], + 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) + for server in servers: + stop_radix_server(server) + 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("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() + 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..d82e2e417 --- /dev/null +++ b/tests/radixshmem/test_e2e_radix_shmem.py @@ -0,0 +1,264 @@ +"""End-to-end test of FLEXKV_ENABLE_RADIXSHMEM=1 on one node: one or two DP scheduler +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; 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 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. + * 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, + start_radix_server, + stop_radix_server, + sweep_radix_files, + wait_kv_manager_ready, + write_pattern, +) + +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, 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. + # 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", + "FLEXKV_RADIXSHMEM_SERVER_NAME": server_name, + "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_server_name = server_name + 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); 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) + 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_") + name = f"/{server_id}" + # 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() + procs = [ + ctx.Process(target=_dp_proc, + args=(dp, dp_size, server_id, name, 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) + stop_radix_server(server) + 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("radixshmem") + 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..98a824d89 --- /dev/null +++ b/tests/radixshmem/test_radix_shmem_engine.py @@ -0,0 +1,1977 @@ +"""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 `radixshmem` is missing. + +Four parts, one file: + + Part 1 — `CacheEngineRadixShmem` semantics against an in-process + `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 + 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 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; + 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 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. +""" +from __future__ import annotations + +import contextlib +import copy +import dataclasses +import glob +import importlib.util +import multiprocessing as mp +import os +import shutil +import socket +import subprocess +import sys +import tempfile +import threading +import time +import traceback +from dataclasses import dataclass +from types import SimpleNamespace + +import numpy as np +import pytest + +try: + import radixshmem +except ImportError as exc: + # 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"radixshmem unusable ({exc}); rebuild the extension", + allow_module_level=True) + +for _name in ("RadixServer", "ServerConfig", "Geometry", "RadixClient"): + if not hasattr(radixshmem, _name): + pytest.skip(f"radixshmem 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 +StagedRadixInsert = _engine_mod.StagedRadixInsert + +# 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.transfer import TransferType # noqa: E402 +from flexkv.server import shm_radix_bootstrap as bootstrap # noqa: E402 + +FULL = radixshmem.ComponentType.FULL +_SWA = radixshmem.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_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 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) + data_bytes, swa_ratio = _server_budget(blocks, swa_slots, slot_bytes) + cfg = radixshmem.ServerConfig(name=name, data_bytes=data_bytes, swa_ratio=swa_ratio, + slot_align=align, prefault=False) + 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) + return cfg, geo + + +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): + """Start the server: ``waiting`` until a client brings the geometry.""" + _sweep_region(cfg.name, cfg.resolved_data_name) + server = radixshmem.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, geo = _server_config(name, blocks, tokens_per_block) + server = self.server(cfg) + engine = self.engine(name, geometry=geo, 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()() + + +@pytest.fixture +def env(): + 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() + bootstrap.READY_TIMEOUT_S = saved + + +# ============================================================================= +# 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_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_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. + + 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_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(). + + 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, geo = _server_config(name, blocks, tokens_per_block, + swa_slots=swa_slots, window_blocks=window_blocks) + server = env.server(cfg) + engine = env.engine( + 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 + + +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. 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 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_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_radixshmem().to_dict() + 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_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_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 + 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 + 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) + 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(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 + client = bootstrap.attach_radix_client(name, geometry=geo, timeout_s=60) + try: + 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; 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, "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( + 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"): + bootstrap.check_geometry( + client, bootstrap.expected_geometry(model_config, cache_config), "test") + # ...and a second client bringing another geometry is refused by the server + 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() + + +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(radixshmem.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: + 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: + client.close() + with pytest.raises(TimeoutError, match="radix-server --name"): + 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(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(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) + 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(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(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(radixshmem, "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 + 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(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) + 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 +# 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 + + +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) + with pytest.raises(ValueError, match="FLEXKV_RADIXSHMEM_SERVER_NAME"): + bootstrap.radix_server_name() + + +# ============================================================================= +# 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 `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): + 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"] == bootstrap.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 = 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)) + 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_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.""" + 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 + + 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 + GLOBAL_CONFIG_FROM_ENV.radixshmem_server_name = server_name + + 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) + # 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 = 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 = 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" + 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) + + +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 == [radixshmem.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 +# radixshmem 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 radixshmem import _data + if not hasattr(_data, "DataPlaneRegistry"): + pytest.skip("radixshmem built without etcd (no DataPlaneRegistry)") + except ImportError: + pytest.skip("radixshmem 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: + # 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(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 = 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=radixshmem.ClusterConfig(**cluster_kwargs), + ) + 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=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: + 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"/radixshmem_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"])) 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_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(