Skip to content

Commit 42ddf58

Browse files
committed
support layerwise
1 parent 2550a9c commit 42ddf58

5 files changed

Lines changed: 319 additions & 26 deletions

File tree

corelib/recsys_kvcache_manager/recsys_kvcache_manager/flex_kvcache_manager.py

Lines changed: 98 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,10 @@
3939
)
4040
from .kvcache_metadata import KVCacheMetadata
4141
from .kvcache_utils import KVIndexMeta, KVLookupResult
42+
from .flexkv_layerwise import (
43+
DEFAULT_LAYERWISE_EVENTFD_SOCKET,
44+
FlexKVLayerwiseEventfdSender,
45+
)
4246

4347

4448
@dataclass
@@ -56,6 +60,12 @@ class _FlexKVOnloadHandle:
5660
task_ids: List[int]
5761
uids: torch.Tensor
5862
slot_mappings: List[torch.Tensor]
63+
layer_eventfds: Optional[List[int]] = None
64+
65+
def wait_layer(self, layer_idx: int) -> None:
66+
if self.layer_eventfds is None:
67+
return
68+
os.read(self.layer_eventfds[layer_idx], 8)
5969

6070

6171
@dataclass
@@ -79,7 +89,7 @@ def __post_init__(self) -> None:
7989
[
8090
self.num_layer,
8191
self.num_block,
82-
self._kv_dim,
92+
self.kv_dim,
8393
self.tokens_per_block,
8494
self.num_head,
8595
self.head_size,
@@ -113,6 +123,10 @@ def __init__(
113123
enable_mps: bool = False,
114124
hostkv_wait_timeout_ms: int = 0,
115125
host_kvstorage_fail_policy: str = "fail_open",
126+
config_path: Optional[str] = None,
127+
enable_layerwise: Optional[bool] = None,
128+
layerwise_eventfd_socket: Optional[str] = None,
129+
layerwise_counter_id: int = 0,
116130
) -> None:
117131
self.mode = mode
118132
self.server_addr = server_addr
@@ -128,6 +142,17 @@ def __init__(
128142
self.hostkv_wait_timeout_ms = int(hostkv_wait_timeout_ms)
129143
self.host_kvstorage_fail_policy = host_kvstorage_fail_policy
130144
self.enable_mps = bool(enable_mps)
145+
self.as_batch = bool(as_batch)
146+
self.config_path = config_path or ""
147+
self.enable_layerwise = enable_layerwise
148+
self.layerwise_eventfd_socket = (
149+
layerwise_eventfd_socket
150+
or os.environ.get(
151+
"FLEXKV_LAYERWISE_EVENTFD_SOCKET",
152+
DEFAULT_LAYERWISE_EVENTFD_SOCKET,
153+
)
154+
)
155+
self.layerwise_counter_id = int(layerwise_counter_id)
131156
self.backend_name = "flexkv"
132157

133158
self._gpu_cache_table_list: Optional[List[torch.Tensor]] = None
@@ -139,6 +164,7 @@ def __init__(
139164
self._adapter = FlexKVClientAdapter(mode, server_addr, server_port)
140165
self._client = None
141166
self._ready = False
167+
self._layerwise_eventfd_sender: Optional[FlexKVLayerwiseEventfdSender] = None
142168

143169
def register_gpu_cache_tables(self, cache_table_list: List[torch.Tensor]) -> None:
144170
assert (
@@ -185,7 +211,7 @@ def _init_client(self) -> None:
185211
if self._client is not None:
186212
return
187213
try:
188-
from flexkv.common.config import CacheConfig, ModelConfig
214+
from flexkv.common.config import CacheConfig, GLOBAL_CONFIG_FROM_ENV, ModelConfig
189215
from flexkv.kvmanager import KVManager
190216
except Exception as e:
191217
raise RuntimeError(f"FlexKV SDK import failed: {e}") from e
@@ -208,6 +234,56 @@ def _init_client(self) -> None:
208234
if self.num_tmp_cpu_blocks > 0:
209235
cache_cfg_kwargs["num_tmp_cpu_blocks"] = self.num_tmp_cpu_blocks
210236
cache_cfg = CacheConfig(**cache_cfg_kwargs)
237+
if self.config_path:
238+
try:
239+
from flexkv.common.config import (
240+
load_user_config_from_file,
241+
update_default_config_from_user_config,
242+
)
243+
except Exception as e:
244+
raise RuntimeError(f"FlexKV config import failed: {e}") from e
245+
246+
user_cfg = load_user_config_from_file(self.config_path)
247+
try:
248+
from flexkv.common.config import RankInfo
249+
250+
rank_info_or_model_cfg = RankInfo(model_config=model_cfg)
251+
except ImportError:
252+
rank_info_or_model_cfg = model_cfg
253+
update_default_config_from_user_config(
254+
rank_info_or_model_cfg, cache_cfg, user_cfg
255+
)
256+
if self.enable_layerwise is None:
257+
enable_layerwise_env = os.environ.get("RECSYS_FLEXKV_ENABLE_LAYERWISE", "0")
258+
self.enable_layerwise = enable_layerwise_env.strip().lower() in {
259+
"1",
260+
"true",
261+
"yes",
262+
"on",
263+
}
264+
elif isinstance(self.enable_layerwise, str):
265+
self.enable_layerwise = self.enable_layerwise.strip().lower() in {
266+
"1",
267+
"true",
268+
"yes",
269+
"on",
270+
}
271+
else:
272+
self.enable_layerwise = bool(self.enable_layerwise)
273+
GLOBAL_CONFIG_FROM_ENV.enable_layerwise_transfer = bool(self.enable_layerwise)
274+
# FlexKV transfer workers are spawned in child processes, so mirror the
275+
# in-process config override into env before KVManager starts them.
276+
os.environ["FLEXKV_ENABLE_LAYERWISE_TRANSFER"] = (
277+
"1" if self.enable_layerwise else "0"
278+
)
279+
if self.enable_layerwise:
280+
os.environ["FLEXKV_LAYERWISE_EVENTFD_SOCKET"] = self.layerwise_eventfd_socket
281+
if self._layerwise_eventfd_sender is None:
282+
self._layerwise_eventfd_sender = FlexKVLayerwiseEventfdSender(
283+
num_layers=self.num_layers,
284+
socket_path=self.layerwise_eventfd_socket,
285+
)
286+
self._layerwise_eventfd_sender.start()
211287
self._client = KVManager(
212288
model_config=model_cfg,
213289
cache_config=cache_cfg,
@@ -373,12 +449,18 @@ def onboard_kvcache_launch(
373449
task_ids=onboard_task_ids,
374450
uids=torch.tensor(onboard_uids, dtype=torch.int64),
375451
slot_mappings=onboard_slot_mappings,
452+
layer_eventfds=self._layerwise_eventfd_sender.layer_eventfds(
453+
self.layerwise_counter_id
454+
)
455+
if self.enable_layerwise and self._layerwise_eventfd_sender is not None
456+
else None,
376457
)
377458
onload_task_handle = HostKVTaskHandle(
378459
backend="flexkv",
379460
user_ids=onload_handle.uids,
380461
handle=onload_handle,
381462
status=HostKVTaskStatus.LAUNCHED,
463+
is_layerwise=bool(self.enable_layerwise),
382464
metadata={
383465
"onboard_start_indices": torch.tensor(
384466
onboard_start_indices, dtype=torch.int32
@@ -387,7 +469,20 @@ def onboard_kvcache_launch(
387469
},
388470
)
389471

390-
self._client.launch(onload_handle.task_ids, onload_handle.slot_mappings)
472+
use_batch = self.as_batch and len(onboard_task_ids) > 1
473+
launch_kwargs = {"as_batch": use_batch}
474+
if self.enable_layerwise:
475+
launch_kwargs.update(
476+
{
477+
"layerwise_transfer": True,
478+
"counter_id": self.layerwise_counter_id,
479+
}
480+
)
481+
self._client.launch(
482+
onload_handle.task_ids,
483+
onload_handle.slot_mappings,
484+
**launch_kwargs,
485+
)
391486
return onload_task_handle
392487

393488
def onboard_kvcache_wait(self, task_handle: HostKVTaskHandle) -> HostKVWaitResult:
Lines changed: 127 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,127 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
#
4+
# Licensed under the Apache License, Version 2.0 (the "License");
5+
# you may not use this file except in compliance with the License.
6+
# You may obtain a copy of the License at
7+
#
8+
# http://www.apache.org/licenses/LICENSE-2.0
9+
#
10+
# Unless required by applicable law or agreed to in writing, software
11+
# distributed under the License is distributed on an "AS IS" BASIS,
12+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
# See the License for the specific language governing permissions and
14+
# limitations under the License.
15+
16+
import os
17+
import socket
18+
import struct
19+
import threading
20+
import time
21+
from array import array
22+
from ctypes import CDLL, get_errno
23+
from typing import List, Optional
24+
25+
26+
DEFAULT_LAYERWISE_EVENTFD_SOCKET = "/tmp/flexkv_layerwise_eventfd.sock"
27+
28+
29+
class FlexKVLayerwiseEventfdSender:
30+
"""Creates mock layerwise eventfds and sends them to FlexKV's worker."""
31+
32+
def __init__(
33+
self,
34+
num_layers: int,
35+
socket_path: str,
36+
num_counters: int = 3,
37+
timeout_s: float = 180.0,
38+
) -> None:
39+
self.num_layers = int(num_layers)
40+
self.socket_path = socket_path
41+
self.num_counters = int(num_counters)
42+
self.timeout_s = float(timeout_s)
43+
self._eventfds: Optional[List[List[int]]] = None
44+
self._thread: Optional[threading.Thread] = None
45+
46+
@staticmethod
47+
def _create_eventfd() -> int:
48+
if hasattr(os, "eventfd"):
49+
return os.eventfd(0, 0)
50+
libc = CDLL(None, use_errno=True)
51+
fd = libc.eventfd(0, 0)
52+
if fd < 0:
53+
raise OSError(get_errno(), "eventfd creation failed")
54+
return int(fd)
55+
56+
def create_eventfds(self) -> List[List[int]]:
57+
if self._eventfds is None:
58+
# FlexKV layerwise worker expects counter sets for triple buffering.
59+
self._eventfds = [
60+
[self._create_eventfd() for _ in range(self.num_layers)]
61+
for _ in range(self.num_counters)
62+
]
63+
return self._eventfds
64+
65+
def layer_eventfds(self, counter_id: int) -> List[int]:
66+
eventfds = self.create_eventfds()
67+
if counter_id < 0 or counter_id >= len(eventfds):
68+
raise ValueError(
69+
f"Invalid layerwise counter_id={counter_id}, "
70+
f"expected [0, {len(eventfds)})"
71+
)
72+
return eventfds[counter_id]
73+
74+
def start(self) -> None:
75+
if self._thread is not None:
76+
return
77+
self.create_eventfds()
78+
self._thread = threading.Thread(
79+
target=self._send_eventfds,
80+
name="flexkv-layerwise-eventfd-sender",
81+
daemon=True,
82+
)
83+
self._thread.start()
84+
85+
def _send_eventfds(self) -> None:
86+
eventfds = self.create_eventfds()
87+
deadline = time.time() + self.timeout_s
88+
last_error: Optional[Exception] = None
89+
while time.time() < deadline:
90+
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
91+
try:
92+
sock.connect(self.socket_path)
93+
metadata = struct.pack(
94+
"iiii",
95+
0,
96+
1,
97+
self.num_layers,
98+
len(eventfds),
99+
)
100+
sock.sendall(metadata)
101+
for counter_id, fds in enumerate(eventfds):
102+
fd_array = array("i", fds)
103+
sock.sendmsg(
104+
[struct.pack("i", counter_id)],
105+
[
106+
(
107+
socket.SOL_SOCKET,
108+
socket.SCM_RIGHTS,
109+
fd_array.tobytes(),
110+
)
111+
],
112+
)
113+
ack = sock.recv(1)
114+
if ack != b"\x01":
115+
raise RuntimeError(
116+
f"FlexKV layerwise eventfd receiver returned ack={ack!r}"
117+
)
118+
return
119+
except (FileNotFoundError, ConnectionRefusedError, socket.timeout) as e:
120+
last_error = e
121+
time.sleep(0.05)
122+
finally:
123+
sock.close()
124+
raise RuntimeError(
125+
"Timed out sending mock layerwise eventfds to FlexKV "
126+
f"socket {self.socket_path}: {last_error}"
127+
)

corelib/recsys_kvcache_manager/recsys_kvcache_manager/kvcache_manager.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -327,6 +327,32 @@ def _build_host_kvstorage_manager_from_config(
327327
}
328328
else:
329329
flexkv_enable_mps = bool(flexkv_enable_mps_raw)
330+
flexkv_as_batch_raw = extra.get("flexkv_as_batch", 1)
331+
if isinstance(flexkv_as_batch_raw, str):
332+
flexkv_as_batch = flexkv_as_batch_raw.strip().lower() in {
333+
"1",
334+
"true",
335+
"yes",
336+
"on",
337+
}
338+
else:
339+
flexkv_as_batch = bool(flexkv_as_batch_raw)
340+
flexkv_enable_layerwise = extra.get("flexkv_enable_layerwise", None)
341+
if isinstance(flexkv_enable_layerwise, str):
342+
flexkv_enable_layerwise = flexkv_enable_layerwise.strip().lower() in {
343+
"1",
344+
"true",
345+
"yes",
346+
"on",
347+
}
348+
elif flexkv_enable_layerwise is not None:
349+
flexkv_enable_layerwise = bool(flexkv_enable_layerwise)
350+
flexkv_layerwise_eventfd_socket = extra.get(
351+
"flexkv_layerwise_eventfd_socket", None
352+
)
353+
flexkv_layerwise_counter_id = int(
354+
extra.get("flexkv_layerwise_counter_id", 0)
355+
)
330356

331357
return FlexKVStorageManager(
332358
mode=flexkv_mode,
@@ -343,6 +369,10 @@ def _build_host_kvstorage_manager_from_config(
343369
enable_mps=flexkv_enable_mps,
344370
host_kvstorage_fail_policy=flexkv_host_kvstorage_fail_policy,
345371
hostkv_wait_timeout_ms=int(kvcache_config.offload_timeout_ms),
372+
config_path=flexkv_config_path,
373+
enable_layerwise=flexkv_enable_layerwise,
374+
layerwise_eventfd_socket=flexkv_layerwise_eventfd_socket,
375+
layerwise_counter_id=flexkv_layerwise_counter_id,
346376
)
347377
else:
348378
raise NotImplementedError(

0 commit comments

Comments
 (0)