Expose max total num tokens from Runtime & Engine API (#2092)
This commit is contained in:
committed by
GitHub
parent
72f87b723b
commit
c35cd1f8c7
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user