3939)
4040from .kvcache_metadata import KVCacheMetadata
4141from .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 :
0 commit comments