Skip to content

Commit e31c5ef

Browse files
committed
use FlexKV main + add test result
1 parent 2b32863 commit e31c5ef

8 files changed

Lines changed: 492 additions & 1058 deletions

File tree

corelib/recsys_kvcache_manager/recsys_kvcache_manager/default_kvcache_backend.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -313,6 +313,22 @@ def _build_host_kvstorage_manager_from_config(
313313
}
314314
else:
315315
flexkv_as_batch = bool(flexkv_as_batch_raw)
316+
flexkv_enable_layerwise = extra.get("flexkv_enable_layerwise", None)
317+
if isinstance(flexkv_enable_layerwise, str):
318+
flexkv_enable_layerwise = flexkv_enable_layerwise.strip().lower() in {
319+
"1",
320+
"true",
321+
"yes",
322+
"on",
323+
}
324+
elif flexkv_enable_layerwise is not None:
325+
flexkv_enable_layerwise = bool(flexkv_enable_layerwise)
326+
flexkv_layerwise_eventfd_socket = extra.get(
327+
"flexkv_layerwise_eventfd_socket", None
328+
)
329+
flexkv_layerwise_counter_id = int(
330+
extra.get("flexkv_layerwise_counter_id", 0)
331+
)
316332

317333
return FlexKVStorage(
318334
mode=flexkv_mode,
@@ -331,6 +347,9 @@ def _build_host_kvstorage_manager_from_config(
331347
host_kvstorage_fail_policy=flexkv_host_kvstorage_fail_policy,
332348
hostkv_wait_timeout_ms=int(kvcache_config.offload_timeout_ms),
333349
config_path=flexkv_config_path,
350+
enable_layerwise=flexkv_enable_layerwise,
351+
layerwise_eventfd_socket=flexkv_layerwise_eventfd_socket,
352+
layerwise_counter_id=flexkv_layerwise_counter_id,
334353
)
335354
else:
336355
raise NotImplementedError(

corelib/recsys_kvcache_manager/recsys_kvcache_manager/flex_kvcache_manager.py

Lines changed: 31 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121

2222
import numpy as np
2323
import torch
24+
from flexkv.common.config import LayerGroupSpec
2425
from flexkv.common.storage import KVCacheLayout, KVCacheLayoutType
2526
from flexkv.server.client import KVTPClient
2627

@@ -31,6 +32,10 @@
3132

3233
KVResponse = Any # type: ignore
3334

35+
from .flexkv_layerwise import (
36+
FlexKVLayerwiseEventfdSender,
37+
create_layerwise_eventfd_socket_path,
38+
)
3439
from .host_kvstorage_manager import (
3540
HostKVStorageBase,
3641
HostKVTaskHandle,
@@ -39,10 +44,6 @@
3944
)
4045
from .kvcache_metadata import KVCacheMetadata
4146
from .kvcache_utils import KVIndexMeta, KVLookupResult
42-
from .flexkv_layerwise import (
43-
DEFAULT_LAYERWISE_EVENTFD_SOCKET,
44-
FlexKVLayerwiseEventfdSender,
45-
)
4647

4748

4849
@dataclass
@@ -148,10 +149,8 @@ def __init__(
148149
self.enable_layerwise = enable_layerwise
149150
self.layerwise_eventfd_socket = (
150151
layerwise_eventfd_socket
151-
or os.environ.get(
152-
"FLEXKV_LAYERWISE_EVENTFD_SOCKET",
153-
DEFAULT_LAYERWISE_EVENTFD_SOCKET,
154-
)
152+
or os.environ.get("FLEXKV_LAYERWISE_EVENTFD_SOCKET")
153+
or create_layerwise_eventfd_socket_path()
155154
)
156155
self.layerwise_counter_id = int(layerwise_counter_id)
157156
self.backend_name = "flexkv"
@@ -190,16 +189,38 @@ def register_gpu_cache_tables(self, cache_table_list: List[torch.Tensor]) -> Non
190189
tokens_per_block=int(first_table.shape[2]),
191190
num_head=int(first_table.shape[3]),
192191
head_size=int(first_table.shape[4]),
193-
is_mla=False,
194192
)
195193
tp_client = KVTPClient(
196194
gpu_register_port=self._gpu_register_port,
197195
dp_client_id=0,
198196
device_id=device_id,
199197
)
198+
register_kwargs: Dict[str, Any] = {}
199+
if self.enable_layerwise:
200+
# FlexKV main's multi-group path is also its public representation
201+
# for a uniform cache with one member per layer. Registering HSTU as
202+
# one group keeps the layout unchanged while avoiding assumptions in
203+
# the legacy single-group layerwise constructor.
204+
register_kwargs.update(
205+
layer_groups=[
206+
LayerGroupSpec(
207+
num_layers=gpu_layout.num_layer,
208+
num_kv_heads=gpu_layout.num_head,
209+
head_size=gpu_layout.head_size,
210+
layer_indices=list(range(gpu_layout.num_layer)),
211+
dtype=self.dtype,
212+
)
213+
],
214+
gpu_layouts=[gpu_layout],
215+
handles_per_group=[self._gpu_cache_table_list],
216+
)
200217
tp_client.register_to_server(
201-
kv_caches=self._gpu_cache_table_list, kv_layout=gpu_layout
218+
kv_caches=self._gpu_cache_table_list,
219+
kv_layout=gpu_layout,
220+
**register_kwargs,
202221
)
222+
if self.enable_layerwise and self._layerwise_eventfd_sender is not None:
223+
self._layerwise_eventfd_sender.wait_until_ready()
203224
self._registered = True
204225

205226
# Client becomes operational only after transfer manager is ready.
@@ -559,9 +580,6 @@ def offload_kvcache_launch(
559580
else self._build_slot_mappings(kvcache_metadata)
560581
)
561582

562-
# Keep offload on the direct put_async path while validating layerwise.
563-
# This matches the known-good 2e259b0 behavior and isolates PR418's
564-
# put_match/as_batch offload optimization from PR428 layerwise onboard.
565583
use_batch = False
566584
task_ids: List[int] = []
567585
batch_slot_mappings: List[torch.Tensor] = []

corelib/recsys_kvcache_manager/recsys_kvcache_manager/flexkv_layerwise.py

Lines changed: 83 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,9 @@
1515

1616
import os
1717
import socket
18+
import stat
1819
import struct
20+
import tempfile
1921
import threading
2022
import time
2123
from array import array
@@ -26,8 +28,16 @@
2628
DEFAULT_LAYERWISE_EVENTFD_SOCKET = "/tmp/flexkv_layerwise_eventfd.sock"
2729

2830

31+
def create_layerwise_eventfd_socket_path() -> str:
32+
socket_directory = tempfile.mkdtemp(
33+
prefix=f"recsys-flexkv-layerwise-{os.geteuid()}-"
34+
)
35+
os.chmod(socket_directory, 0o700)
36+
return os.path.join(socket_directory, "eventfd.sock")
37+
38+
2939
class FlexKVLayerwiseEventfdSender:
30-
"""Creates mock layerwise eventfds and sends them to FlexKV's worker."""
40+
"""Creates layerwise eventfds and sends them to FlexKV's worker."""
3141

3242
def __init__(
3343
self,
@@ -42,6 +52,8 @@ def __init__(
4252
self.timeout_s = float(timeout_s)
4353
self._eventfds: Optional[List[List[int]]] = None
4454
self._thread: Optional[threading.Thread] = None
55+
self._handoff_done = threading.Event()
56+
self._handoff_error: Optional[BaseException] = None
4557

4658
@staticmethod
4759
def _create_eventfd() -> int:
@@ -76,20 +88,87 @@ def start(self) -> None:
7688
return
7789
self.create_eventfds()
7890
self._thread = threading.Thread(
79-
target=self._send_eventfds,
91+
target=self._run_sender,
8092
name="flexkv-layerwise-eventfd-sender",
8193
daemon=True,
8294
)
8395
self._thread.start()
8496

97+
def wait_until_ready(self, timeout_s: Optional[float] = None) -> None:
98+
timeout = self.timeout_s + 1.0 if timeout_s is None else float(timeout_s)
99+
if not self._handoff_done.wait(timeout):
100+
raise TimeoutError(
101+
"Timed out waiting for FlexKV layerwise eventfd handoff "
102+
f"on socket {self.socket_path}"
103+
)
104+
if self._handoff_error is not None:
105+
raise RuntimeError(
106+
"FlexKV layerwise eventfd handoff failed"
107+
) from self._handoff_error
108+
109+
def _run_sender(self) -> None:
110+
try:
111+
self._send_eventfds()
112+
except BaseException as error:
113+
self._handoff_error = error
114+
finally:
115+
self._handoff_done.set()
116+
117+
@staticmethod
118+
def _is_process_descendant(process_id: int, ancestor_id: int) -> bool:
119+
current_id = int(process_id)
120+
ancestor_id = int(ancestor_id)
121+
visited = set()
122+
while current_id > 1 and current_id not in visited:
123+
if current_id == ancestor_id:
124+
return True
125+
visited.add(current_id)
126+
try:
127+
with open(f"/proc/{current_id}/stat", encoding="utf-8") as stat_file:
128+
process_stat = stat_file.read()
129+
fields_after_name = process_stat[process_stat.rfind(")") + 1 :].split()
130+
current_id = int(fields_after_name[1])
131+
except (OSError, IndexError, ValueError):
132+
return False
133+
return current_id == ancestor_id
134+
135+
def _authenticate_peer(self, sock: socket.socket) -> None:
136+
socket_stat = os.stat(self.socket_path, follow_symlinks=False)
137+
if not stat.S_ISSOCK(socket_stat.st_mode):
138+
raise PermissionError(
139+
f"Layerwise eventfd path is not a socket: {self.socket_path}"
140+
)
141+
if socket_stat.st_uid != os.geteuid():
142+
raise PermissionError(
143+
"FlexKV layerwise socket owner does not match the current user"
144+
)
145+
146+
if not hasattr(socket, "SO_PEERCRED"):
147+
raise RuntimeError("SO_PEERCRED is required for layerwise eventfd handoff")
148+
credentials_size = struct.calcsize("3i")
149+
peer_pid, peer_uid, peer_gid = struct.unpack(
150+
"3i",
151+
sock.getsockopt(socket.SOL_SOCKET, socket.SO_PEERCRED, credentials_size),
152+
)
153+
if peer_uid != os.geteuid() or peer_gid != os.getegid():
154+
raise PermissionError(
155+
"FlexKV layerwise peer credentials do not match the current process"
156+
)
157+
if not self._is_process_descendant(peer_pid, os.getpid()):
158+
raise PermissionError(
159+
f"FlexKV layerwise peer pid {peer_pid} is outside the current process tree"
160+
)
161+
85162
def _send_eventfds(self) -> None:
86163
eventfds = self.create_eventfds()
87164
deadline = time.time() + self.timeout_s
88165
last_error: Optional[Exception] = None
89166
while time.time() < deadline:
90167
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
91168
try:
169+
sock.settimeout(min(1.0, max(deadline - time.time(), 0.01)))
92170
sock.connect(self.socket_path)
171+
self._authenticate_peer(sock)
93172
metadata = struct.pack(
94173
"iiii",
95174
0,
@@ -116,12 +195,12 @@ def _send_eventfds(self) -> None:
116195
f"FlexKV layerwise eventfd receiver returned ack={ack!r}"
117196
)
118197
return
119-
except (FileNotFoundError, ConnectionRefusedError, socket.timeout) as e:
198+
except (OSError, RuntimeError) as e:
120199
last_error = e
121200
time.sleep(0.05)
122201
finally:
123202
sock.close()
124203
raise RuntimeError(
125-
"Timed out sending mock layerwise eventfds to FlexKV "
204+
"Timed out sending layerwise eventfds to FlexKV "
126205
f"socket {self.socket_path}: {last_error}"
127206
)

corelib/recsys_kvcache_manager/recsys_kvcache_manager/host_kvstorage_manager.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -119,11 +119,14 @@ def onboard_kvcache_launch(
119119
def onboard_kvcache_wait(self, task_handle: HostKVTaskHandle) -> HostKVWaitResult:
120120
...
121121

122-
@abstractmethod
123122
def onboard_kvcache_wait_by_layer(
124123
self, task_handle: HostKVTaskHandle, layer_idx: int
125124
) -> HostKVWaitResult:
126-
...
125+
return HostKVWaitResult(
126+
status=HostKVTaskStatus.SKIPPED,
127+
ready=False,
128+
message="Layerwise onboard wait is not supported by this backend.",
129+
)
127130

128131
@abstractmethod
129132
def offload_kvcache_launch(

corelib/recsys_kvcache_manager/test/test_flexkv.py

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,17 +22,21 @@
2222
from recsys_kvcache_manager.kvcache_manager import KVCacheManager
2323

2424

25-
def _get_flexkv_enable_layerwise() -> str:
26-
return os.environ.get("RECSYS_FLEXKV_ENABLE_LAYERWISE", "")
25+
def _get_flexkv_config_path() -> str:
26+
return os.environ.get("RECSYS_FLEXKV_CONFIG_PATH", "")
2727

2828

2929
def create_testing_kvcache_manager() -> KVCacheManager:
30-
flexkv_enable_layerwise = _get_flexkv_enable_layerwise()
30+
flexkv_config_path = _get_flexkv_config_path()
31+
flexkv_enable_layerwise = os.environ.get("RECSYS_FLEXKV_ENABLE_LAYERWISE", "")
3132
extra_configs = {
3233
"flexkv_mode": "direct",
3334
"flexkv_host_kvstorage_fail_policy": "fail_open",
3435
"flexkv_enable_mps": 0,
36+
"flexkv_as_batch": 1,
3537
}
38+
if flexkv_config_path:
39+
extra_configs["flexkv_config_path"] = flexkv_config_path
3640
if flexkv_enable_layerwise:
3741
extra_configs["flexkv_enable_layerwise"] = flexkv_enable_layerwise
3842

@@ -68,15 +72,27 @@ def create_testing_kvcache_manager() -> KVCacheManager:
6872
)
6973
print(f"[TEST] KVCache GPU Memory Usage: {gpu_gib} GiB.")
7074
print(f"[TEST] KVCache Host Memory Usage: {host_gib} GiB.")
75+
if flexkv_config_path:
76+
print(f"[TEST] FlexKV config path: {flexkv_config_path}")
7177
kvcache_mgr = KVCacheManager.from_config(kvcache_config)
7278
flexkv_mgr = kvcache_mgr.host_kvstorage_manager
79+
cache_cfg = flexkv_mgr._client.cache_config
80+
if cache_cfg.enable_ssd:
81+
print(
82+
"[TEST] Created KVCache Manager with FlexKV SSD tier: "
83+
f"num_cpu_blocks={cache_cfg.num_cpu_blocks}, "
84+
f"num_ssd_blocks={cache_cfg.num_ssd_blocks}, "
85+
f"ssd_cache_dir={cache_cfg.ssd_cache_dir}, "
86+
f"enable_gds={cache_cfg.enable_gds}"
87+
)
88+
else:
89+
print("[TEST] Created KVCache Manager with FlexKV CPU tier only")
7390
if flexkv_mgr.enable_layerwise:
7491
print(
7592
"[TEST] FlexKV layerwise transfer enabled: "
7693
f"eventfd_socket={flexkv_mgr.layerwise_eventfd_socket}, "
7794
f"counter_id={flexkv_mgr.layerwise_counter_id}"
7895
)
79-
print("[TEST] Created KVCache Manager with FlexKV CPU tier only")
8096
return kvcache_mgr
8197

8298

0 commit comments

Comments
 (0)