1+ import os
12import tempfile
23from multiprocessing import Process
34import argparse
89import torch
910
1011from flexkv .server .client import KVDPClient , KVTPClient
11- from flexkv .server .server import KVServer
12+ from flexkv .server .server import KVServer , SchedulerServer
1213from flexkv .common .config import ModelConfig , CacheConfig
1314from flexkv .common .storage import KVCacheLayoutType , KVCacheLayout
1415from flexkv .common .debug import flexkv_logger
@@ -58,20 +59,7 @@ def run_tp_client(dp_client_id, tp_rank, server_recv_port, model_config, cache_c
5859 while True :
5960 time .sleep (1 )
6061
61- def shutdown (server_process , dp_client , tp_client_processes , server_recv_port ):
62- """Shutdown all processes"""
63- try :
64- # Send a shutdown request to the server
65- from flexkv .server .request import ShutdownRequest
66- shutdown_request = ShutdownRequest (dp_client_id = dp_client .dp_client_id )
67- dp_client .send_to_server .send_pyobj (shutdown_request )
68-
69- # Wait a bit for graceful shutdown
70- time .sleep (3 )
71- except Exception as e :
72- print (f"Error sending shutdown request: { e } " )
73-
74- # Terminate tp_client processes
62+ def shutdown_tp_client (tp_client_processes ):
7563 for tp_process in tp_client_processes :
7664 if tp_process .is_alive ():
7765 tp_process .terminate ()
@@ -81,24 +69,95 @@ def shutdown(server_process, dp_client, tp_client_processes, server_recv_port):
8169 tp_process .kill ()
8270 tp_process .join (timeout = 2 )
8371
84- # Terminate server process
85- if server_process .is_alive ():
86- server_process .terminate ()
87- server_process .join (timeout = 10 )
88- if server_process .is_alive ():
89- print (f"Force killing server process { server_process .pid } " )
90- server_process .kill ()
91- server_process .join (timeout = 5 )
92-
93- # Clean up temporary file
94- import os
95- if server_recv_port .startswith ('ipc://' ):
96- temp_file = server_recv_port [6 :] # Remove 'ipc://' prefix
97- try :
98- if os .path .exists (temp_file ):
99- os .unlink (temp_file )
100- except Exception as e :
101- print (f"Error cleaning up temporary file: { e } " )
72+ class FlexkvWrapper :
73+ def __init__ (self , model_config , cache_config , server_recv_port ):
74+ self .model_config = model_config
75+ self .cache_config = cache_config
76+ self .server_recv_port = server_recv_port
77+
78+ self .use_scheduler_server = model_config .dp_size == 1
79+ if self .use_scheduler_server :
80+ self .launch_scheduler_server ()
81+ else :
82+ self .launch_server ()
83+
84+ def launch_server (self ):
85+ def server_process ():
86+ kvserver = KVServer (self .model_config , self .cache_config , self .server_recv_port )
87+ kvserver .run ()
88+ time .sleep (10 )
89+ self .server_process = Process (
90+ target = server_process ,
91+ daemon = False
92+ )
93+ self .server_process .start ()
94+ time .sleep (5 )
95+ self .dp_client = KVDPClient (self .server_recv_port , self .model_config )
96+
97+ def launch_scheduler_server (self ):
98+ self .scheduler_server = SchedulerServer (self .model_config , self .cache_config , self .server_recv_port )
99+ self .scheduler_server .start_server_thread ()
100+ time .sleep (10 )
101+
102+ @property
103+ def dp_client_id (self ):
104+ if self .use_scheduler_server :
105+ return 0
106+ else :
107+ return self .dp_client .dp_client_id
108+
109+ def put_async (self , token_ids , slot_mapping , token_mask = None ):
110+ if self .use_scheduler_server :
111+ return self .scheduler_server .put_async (token_ids , slot_mapping , token_mask )
112+ else :
113+ return self .dp_client .put_async (token_ids , slot_mapping , token_mask )
114+
115+ def get_async (self , token_ids , slot_mapping , token_mask = None ):
116+ if self .use_scheduler_server :
117+ return self .scheduler_server .get_async (token_ids , slot_mapping , token_mask )
118+ else :
119+ return self .dp_client .get_async (token_ids , slot_mapping , token_mask )
120+
121+ def wait (self , request_ids ):
122+ if self .use_scheduler_server :
123+ return self .scheduler_server .wait (request_ids )
124+ else :
125+ return self .dp_client .wait (request_ids )
126+
127+ def try_wait (self , request_ids ):
128+ if self .use_scheduler_server :
129+ return self .scheduler_server .try_wait (request_ids )
130+ else :
131+ return self .dp_client .try_wait (request_ids )
132+
133+ def shutdown (self ):
134+ if not self .use_scheduler_server :
135+ try :
136+ # Send a shutdown request to the server
137+ from flexkv .server .request import ShutdownRequest
138+ shutdown_request = ShutdownRequest (dp_client_id = self .dp_client_id )
139+ self .dp_client .send_to_server .send_pyobj (shutdown_request )
140+
141+ # Wait a bit for graceful shutdown
142+ time .sleep (3 )
143+ except Exception as e :
144+ print (f"Error sending shutdown request: { e } " )
145+ if self .server_process .is_alive ():
146+ self .server_process .terminate ()
147+ self .server_process .join (timeout = 10 )
148+ if self .server_process .is_alive ():
149+ print (f"Force killing server process { self .server_process .pid } " )
150+ self .server_process .kill ()
151+ self .server_process .join (timeout = 5 )
152+ if self .server_recv_port .startswith ('ipc://' ):
153+ temp_file = self .server_recv_port [6 :] # Remove 'ipc://' prefix
154+ try :
155+ if os .path .exists (temp_file ):
156+ os .unlink (temp_file )
157+ except Exception as e :
158+ print (f"Error cleaning up temporary file: { e } " )
159+ else :
160+ self .scheduler_server .shutdown ()
102161
103162def benchmark_kvmanager (model_config , cache_config , benchmark_config , server_recv_port ):
104163 if model_config .tp_size * model_config .dp_size > torch .cuda .device_count ():
@@ -107,14 +166,8 @@ def benchmark_kvmanager(model_config, cache_config, benchmark_config, server_rec
107166 print (f"{ model_config = } " )
108167 print (f"{ cache_config = } " )
109168 print (f"{ benchmark_config = } " )
110- server_process = Process (
111- target = run_server ,
112- args = (model_config , cache_config , server_recv_port ),
113- daemon = False
114- )
115- server_process .start ()
116- time .sleep (5 )
117- dp_client = KVDPClient (server_recv_port , model_config )
169+ flexkv_wrapper = FlexkvWrapper (model_config , cache_config , server_recv_port )
170+
118171 tp_client_processes = []
119172
120173 sequence_length = benchmark_config .sequence_length
@@ -125,7 +178,7 @@ def benchmark_kvmanager(model_config, cache_config, benchmark_config, server_rec
125178 for tp_rank in range (model_config .tp_size ):
126179 tp_client_process = Process (
127180 target = run_tp_client ,
128- args = (dp_client .dp_client_id , tp_rank , server_recv_port ,
181+ args = (flexkv_wrapper .dp_client_id , tp_rank , server_recv_port ,
129182 model_config , cache_config ),
130183 daemon = True
131184 )
@@ -147,10 +200,10 @@ def benchmark_kvmanager(model_config, cache_config, benchmark_config, server_rec
147200 put_ids = []
148201 if benchmark_config .cache_ratio > 0 :
149202 for i in range (batch_size ):
150- put_ids .append (dp_client .put_async (batch_sequence_tensor [i ][:cache_length ],
203+ put_ids .append (flexkv_wrapper .put_async (batch_sequence_tensor [i ][:cache_length ],
151204 batch_slot_mapping [i ][:cache_length ],
152205 token_mask = None ))
153- put_result = dp_client .wait (put_ids )
206+ put_result = flexkv_wrapper .wait (put_ids )
154207 end_time = time .time ()
155208 time .sleep (1 )
156209 elapsed_time_put = end_time - start_time
@@ -166,10 +219,10 @@ def benchmark_kvmanager(model_config, cache_config, benchmark_config, server_rec
166219 start_time = time .time ()
167220 get_ids = []
168221 for i in range (batch_size ):
169- get_ids .append (dp_client .get_async (batch_sequence_tensor [i ],
222+ get_ids .append (flexkv_wrapper .get_async (batch_sequence_tensor [i ],
170223 batch_slot_mapping [i ],
171224 token_mask = None ))
172- get_result = dp_client .wait (get_ids )
225+ get_result = flexkv_wrapper .wait (get_ids )
173226 end_time = time .time ()
174227 elapsed_time_get = end_time - start_time
175228 cached_tokens = 0
@@ -183,16 +236,18 @@ def benchmark_kvmanager(model_config, cache_config, benchmark_config, server_rec
183236 f"cache_ratio: { cached_tokens * 100 / all_tokens :.2f} %, "
184237 f"time: { elapsed_time_get * 1000 :.2f} ms, bandwidth: { transfer_bandwidth_get :.2f} GB/s" )
185238
186- shutdown (server_process , dp_client , tp_client_processes , server_recv_port )
239+ shutdown_tp_client (tp_client_processes )
240+ flexkv_wrapper .shutdown ()
241+
187242
188243def parse_args ():
189244 parser = argparse .ArgumentParser ()
190245 parser .add_argument ("--config" , type = str , default = "benchmarks/example_config.json" )
191246 # benchmark config
192- parser .add_argument ("--num_layers " , type = int , default = - 1 )
193- parser .add_argument ("--batch_size " , type = int , default = 1 )
194- parser .add_argument ("--sequence_length " , type = int , default = 1024 )
195- parser .add_argument ("--cache_ratio " , type = float , default = 1 )
247+ parser .add_argument ("--num-layers " , type = int , default = - 1 )
248+ parser .add_argument ("--batch-size " , type = int , default = 1 )
249+ parser .add_argument ("--sequence-length " , type = int , default = 1024 )
250+ parser .add_argument ("--cache-ratio " , type = float , default = 1 )
196251 return parser .parse_args ()
197252
198253if __name__ == "__main__" :
0 commit comments