Skip to content

Commit 508a1b1

Browse files
authored
Unify the interface of the flexkv server (#60)
1 parent 067c2b6 commit 508a1b1

3 files changed

Lines changed: 432 additions & 460 deletions

File tree

benchmarks/benchmark_kvmanager.py

Lines changed: 106 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import os
12
import tempfile
23
from multiprocessing import Process
34
import argparse
@@ -8,7 +9,7 @@
89
import torch
910

1011
from flexkv.server.client import KVDPClient, KVTPClient
11-
from flexkv.server.server import KVServer
12+
from flexkv.server.server import KVServer, SchedulerServer
1213
from flexkv.common.config import ModelConfig, CacheConfig
1314
from flexkv.common.storage import KVCacheLayoutType, KVCacheLayout
1415
from 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

103162
def 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

188243
def 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

198253
if __name__ == "__main__":

0 commit comments

Comments
 (0)