Expose max total num tokens from Runtime & Engine API (#2092)

This commit is contained in:
Henry Hyeonmok Ko
2024-11-22 15:10:10 -08:00
committed by GitHub
parent 72f87b723b
commit c35cd1f8c7
4 changed files with 81 additions and 7 deletions

View File

@@ -167,9 +167,12 @@ class DataParallelController:
self.context, zmq.PUSH, port_args.scheduler_input_ipc_name
)
# Wait for model to finish loading
# Wait for model to finish loading and get max token nums
scheduler_info = []
for i in range(len(scheduler_pipe_readers)):
scheduler_pipe_readers[i].recv()
scheduler_info.append(scheduler_pipe_readers[i].recv())
self.max_total_num_tokens = scheduler_info[0]["max_total_num_tokens"]
return send_to
@@ -191,7 +194,10 @@ class DataParallelController:
send_to = get_zmq_socket(
self.context, zmq.PUSH, port_args.scheduler_input_ipc_name
)
reader.recv()
scheduler_info = reader.recv()
self.max_total_num_tokens = scheduler_info["max_total_num_tokens"]
return send_to
def round_robin_scheduler(self, req):
@@ -233,7 +239,9 @@ def run_data_parallel_controller_process(
try:
controller = DataParallelController(server_args, port_args)
pipe_writer.send("ready")
pipe_writer.send(
{"status": "ready", "max_total_num_tokens": controller.max_total_num_tokens}
)
controller.event_loop()
except Exception:
msg = get_exception_traceback()