[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:
Lianmin Zheng
2025-12-09 18:38:38 -08:00
committed by GitHub
co-authored by github-actions[bot]
parent f077436831
commit 4285e99da7
2 changed files with 51 additions and 25 deletions
@@ -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: