Expose max total num tokens from Runtime & Engine API (#2092)
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -1400,7 +1400,9 @@ def run_scheduler_process(
|
||||
|
||||
try:
|
||||
scheduler = Scheduler(server_args, port_args, gpu_id, tp_rank, dp_rank)
|
||||
pipe_writer.send("ready")
|
||||
pipe_writer.send(
|
||||
{"status": "ready", "max_total_num_tokens": scheduler.max_total_num_tokens}
|
||||
)
|
||||
if scheduler.enable_overlap:
|
||||
scheduler.event_loop_overlap()
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user