[Auto Sync] Update data_parallel_controller.py, detokenizer... (20251209) (#14759)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
github-actions[bot]
parent
f077436831
commit
4285e99da7
@@ -21,7 +21,7 @@ import threading
|
|||||||
import time
|
import time
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import List, Optional
|
from typing import Callable, List, Optional
|
||||||
|
|
||||||
import psutil
|
import psutil
|
||||||
import setproctitle
|
import setproctitle
|
||||||
@@ -119,14 +119,19 @@ class DPBudget:
|
|||||||
class DataParallelController:
|
class DataParallelController:
|
||||||
"""A controller that dispatches requests to multiple data parallel workers."""
|
"""A controller that dispatches requests to multiple data parallel workers."""
|
||||||
|
|
||||||
def __init__(self, server_args: ServerArgs, port_args: PortArgs) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
port_args: PortArgs,
|
||||||
|
run_scheduler_process_func: Callable,
|
||||||
|
) -> None:
|
||||||
# Parse args
|
# Parse args
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.port_args = port_args
|
self.port_args = port_args
|
||||||
self.load_balance_method = LoadBalanceMethod.from_str(
|
self.load_balance_method = LoadBalanceMethod.from_str(
|
||||||
server_args.load_balance_method
|
server_args.load_balance_method
|
||||||
)
|
)
|
||||||
self.run_scheduler_process = run_scheduler_process
|
self.run_scheduler_process_func = run_scheduler_process_func
|
||||||
|
|
||||||
# For DP balance
|
# For DP balance
|
||||||
self.global_balance_id = 0
|
self.global_balance_id = 0
|
||||||
@@ -429,7 +434,7 @@ class DataParallelController:
|
|||||||
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
||||||
with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id:
|
with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id:
|
||||||
proc = mp.Process(
|
proc = mp.Process(
|
||||||
target=self.run_scheduler_process,
|
target=self.run_scheduler_process_func,
|
||||||
args=(
|
args=(
|
||||||
server_args,
|
server_args,
|
||||||
rank_port_args,
|
rank_port_args,
|
||||||
@@ -511,7 +516,7 @@ def run_data_parallel_controller_process(
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
port_args: PortArgs,
|
port_args: PortArgs,
|
||||||
pipe_writer,
|
pipe_writer,
|
||||||
data_parallel_controller_class=DataParallelController,
|
run_scheduler_process_func: Callable = run_scheduler_process,
|
||||||
):
|
):
|
||||||
setproctitle.setproctitle("sglang::data_parallel_controller")
|
setproctitle.setproctitle("sglang::data_parallel_controller")
|
||||||
faulthandler.enable()
|
faulthandler.enable()
|
||||||
@@ -529,7 +534,9 @@ def run_data_parallel_controller_process(
|
|||||||
trace_set_thread_info(thread_label)
|
trace_set_thread_info(thread_label)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
controller = data_parallel_controller_class(server_args, port_args)
|
controller = DataParallelController(
|
||||||
|
server_args, port_args, run_scheduler_process_func
|
||||||
|
)
|
||||||
pipe_writer.send(
|
pipe_writer.send(
|
||||||
{
|
{
|
||||||
"status": "ready",
|
"status": "ready",
|
||||||
|
|||||||
@@ -84,6 +84,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
context, zmq.PUSH, port_args.tokenizer_ipc_name, False
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Init tokenizer
|
||||||
if server_args.skip_tokenizer_init:
|
if server_args.skip_tokenizer_init:
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
else:
|
else:
|
||||||
@@ -95,8 +96,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
|
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
|
||||||
self.is_dummy = server_args.load_format == "dummy"
|
self.is_dummy = False
|
||||||
|
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
|
||||||
|
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
|
||||||
|
|
||||||
|
# Init dispatcher
|
||||||
self._request_dispatcher = TypeBasedDispatcher(
|
self._request_dispatcher = TypeBasedDispatcher(
|
||||||
[
|
[
|
||||||
(BatchEmbeddingOutput, self.handle_batch_embedding_out),
|
(BatchEmbeddingOutput, self.handle_batch_embedding_out),
|
||||||
@@ -106,9 +110,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
|
|
||||||
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
|
|
||||||
|
|
||||||
def event_loop(self):
|
def event_loop(self):
|
||||||
"""The event loop that handles requests"""
|
"""The event loop that handles requests"""
|
||||||
while True:
|
while True:
|
||||||
@@ -148,7 +149,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
# If it is embedding model, no detokenization is needed.
|
# If it is embedding model, no detokenization is needed.
|
||||||
return recv_obj
|
return recv_obj
|
||||||
|
|
||||||
def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput):
|
def _decode_batch_token_id_output(self, recv_obj: BatchTokenIDOutput):
|
||||||
bs = len(recv_obj.rids)
|
bs = len(recv_obj.rids)
|
||||||
|
|
||||||
# Initialize decode status
|
# Initialize decode status
|
||||||
@@ -176,8 +177,31 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
)
|
)
|
||||||
surr_ids.append(s.decode_ids[s.surr_offset : s.read_offset])
|
surr_ids.append(s.decode_ids[s.surr_offset : s.read_offset])
|
||||||
|
|
||||||
# TODO(lmzheng): better handle skip_special_tokens/spaces_between_special_tokens per request
|
# Decode token ids to strings
|
||||||
if self.disable_tokenizer_batch_decode:
|
# TODO(lmzheng): handle skip_special_tokens/spaces_between_special_tokens per request
|
||||||
|
if not self.disable_tokenizer_batch_decode:
|
||||||
|
if not self.is_dummy:
|
||||||
|
# Run normal batch decode
|
||||||
|
surr_texts = self.tokenizer.batch_decode(
|
||||||
|
surr_ids,
|
||||||
|
skip_special_tokens=recv_obj.skip_special_tokens[0],
|
||||||
|
spaces_between_special_tokens=recv_obj.spaces_between_special_tokens[
|
||||||
|
0
|
||||||
|
],
|
||||||
|
)
|
||||||
|
read_texts = self.tokenizer.batch_decode(
|
||||||
|
read_ids,
|
||||||
|
skip_special_tokens=recv_obj.skip_special_tokens[0],
|
||||||
|
spaces_between_special_tokens=recv_obj.spaces_between_special_tokens[
|
||||||
|
0
|
||||||
|
],
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# If it is dummy weights, just return dummy strings to prevent potential detokenization edge cases
|
||||||
|
surr_texts = ["dog" for _ in surr_ids]
|
||||||
|
read_texts = ["cat" for _ in read_ids]
|
||||||
|
else:
|
||||||
|
# Do not use batch decode to prevent some detokenization edge cases (e.g., gpt-oss).
|
||||||
surr_texts = [
|
surr_texts = [
|
||||||
self.tokenizer.decode(
|
self.tokenizer.decode(
|
||||||
surr, skip_special_tokens=skip, spaces_between_special_tokens=space
|
surr, skip_special_tokens=skip, spaces_between_special_tokens=space
|
||||||
@@ -198,17 +222,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
recv_obj.spaces_between_special_tokens,
|
recv_obj.spaces_between_special_tokens,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
else:
|
|
||||||
surr_texts = self.tokenizer.batch_decode(
|
|
||||||
surr_ids,
|
|
||||||
skip_special_tokens=recv_obj.skip_special_tokens[0],
|
|
||||||
spaces_between_special_tokens=recv_obj.spaces_between_special_tokens[0],
|
|
||||||
)
|
|
||||||
read_texts = self.tokenizer.batch_decode(
|
|
||||||
read_ids,
|
|
||||||
skip_special_tokens=recv_obj.skip_special_tokens[0],
|
|
||||||
spaces_between_special_tokens=recv_obj.spaces_between_special_tokens[0],
|
|
||||||
)
|
|
||||||
|
|
||||||
# Incremental decoding
|
# Incremental decoding
|
||||||
output_strs = []
|
output_strs = []
|
||||||
@@ -247,6 +260,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
s.sent_offset = len(output_str)
|
s.sent_offset = len(output_str)
|
||||||
output_strs.append(incremental_output)
|
output_strs.append(incremental_output)
|
||||||
|
|
||||||
|
return output_strs
|
||||||
|
|
||||||
|
def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput):
|
||||||
|
output_strs = self._decode_batch_token_id_output(recv_obj)
|
||||||
|
|
||||||
return BatchStrOutput(
|
return BatchStrOutput(
|
||||||
rids=recv_obj.rids,
|
rids=recv_obj.rids,
|
||||||
http_worker_ipcs=recv_obj.http_worker_ipcs,
|
http_worker_ipcs=recv_obj.http_worker_ipcs,
|
||||||
@@ -306,6 +324,7 @@ class LimitedCapacityDict(OrderedDict):
|
|||||||
def run_detokenizer_process(
|
def run_detokenizer_process(
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
port_args: PortArgs,
|
port_args: PortArgs,
|
||||||
|
detokenizer_manager_class=DetokenizerManager,
|
||||||
):
|
):
|
||||||
kill_itself_when_parent_died()
|
kill_itself_when_parent_died()
|
||||||
setproctitle.setproctitle("sglang::detokenizer")
|
setproctitle.setproctitle("sglang::detokenizer")
|
||||||
@@ -313,7 +332,7 @@ def run_detokenizer_process(
|
|||||||
parent_process = psutil.Process().parent()
|
parent_process = psutil.Process().parent()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
manager = DetokenizerManager(server_args, port_args)
|
manager = detokenizer_manager_class(server_args, port_args)
|
||||||
if server_args.tokenizer_worker_num > 1:
|
if server_args.tokenizer_worker_num > 1:
|
||||||
manager.multi_http_worker_event_loop()
|
manager.multi_http_worker_event_loop()
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user