[tokenizer] Support pluggable tokenizer worker class in multi-tokenizer mode (#30630)

Co-authored-by: tanujtiwari1998 <168470992+tanujtiwari1998@users.noreply.github.com>
This commit is contained in:
Lianmin Zheng
2026-07-09 17:31:54 -07:00
committed by GitHub
co-authored by tanujtiwari1998
parent 77b7698cad
commit cfc66e05c5
4 changed files with 63 additions and 3 deletions
+3 -1
View File
@@ -146,6 +146,7 @@ from sglang.srt.managers.multi_tokenizer_mixin import (
MultiTokenizerRouter, MultiTokenizerRouter,
TokenizerWorker, TokenizerWorker,
get_main_process_id, get_main_process_id,
get_tokenizer_worker_class,
read_from_shared_memory, read_from_shared_memory,
write_data_for_multi_tokenizer, write_data_for_multi_tokenizer,
) )
@@ -236,7 +237,8 @@ async def init_multi_tokenizer() -> ServerArgs:
) )
# Launch multi-tokenizer manager process # Launch multi-tokenizer manager process
tokenizer_manager = TokenizerWorker(server_args, port_args) tokenizer_worker_class = get_tokenizer_worker_class(server_args)
tokenizer_manager = tokenizer_worker_class(server_args, port_args)
template_manager = TemplateManager() template_manager = TemplateManager()
template_manager.initialize_templates( template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager, tokenizer_manager=tokenizer_manager,
@@ -29,7 +29,7 @@ import sys
import threading import threading
import zlib import zlib
from multiprocessing import shared_memory from multiprocessing import shared_memory
from typing import TYPE_CHECKING, Any, Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type
import psutil import psutil
import setproctitle import setproctitle
@@ -679,6 +679,19 @@ class TokenizerWorker(TokenizerManager):
self._pause_continue_future = None self._pause_continue_future = None
def get_tokenizer_worker_class(server_args: ServerArgs) -> Type[TokenizerWorker]:
worker_class = server_args.get_tokenizer_worker_class()
if not isinstance(worker_class, type) or not issubclass(
worker_class, TokenizerWorker
):
raise TypeError(
"ServerArgs.get_tokenizer_worker_class() must return a TokenizerWorker "
f"subclass, got {worker_class!r}"
)
return worker_class
async def print_exception_wrapper(func): async def print_exception_wrapper(func):
""" """
Sometimes an asyncio function does not print exception. Sometimes an asyncio function does not print exception.
+5
View File
@@ -6670,6 +6670,11 @@ class ServerArgs:
] ]
return cls(**{attr: getattr(args, attr) for attr in attrs}) return cls(**{attr: getattr(args, attr) for attr in attrs})
def get_tokenizer_worker_class(self):
from sglang.srt.managers.multi_tokenizer_mixin import TokenizerWorker
return TokenizerWorker
def url(self, port: Optional[int] = None): def url(self, port: Optional[int] = None):
scheme = "https" if self.ssl_certfile else "http" scheme = "https" if self.ssl_certfile else "http"
# When binding to all interfaces, use loopback for internal requests. # When binding to all interfaces, use loopback for internal requests.
@@ -6,11 +6,38 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel
maybe_stub_sgl_kernel() maybe_stub_sgl_kernel()
from sglang.srt.managers.io_struct import BatchStrOutput from sglang.srt.managers.io_struct import BatchStrOutput
from sglang.srt.managers.multi_tokenizer_mixin import _handle_output_by_index from sglang.srt.managers.multi_tokenizer_mixin import (
TokenizerWorker,
_handle_output_by_index,
get_tokenizer_worker_class,
)
register_cpu_ci(est_time=5, suite="base-a-test-cpu") register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class CustomTokenizerWorker(TokenizerWorker):
pass
class NotAWorker:
pass
class DefaultServerArgs:
def get_tokenizer_worker_class(self):
return TokenizerWorker
class CustomServerArgs:
def get_tokenizer_worker_class(self):
return CustomTokenizerWorker
class InvalidServerArgs:
def get_tokenizer_worker_class(self):
return NotAWorker
def _make_batch_str_output() -> BatchStrOutput: def _make_batch_str_output() -> BatchStrOutput:
return BatchStrOutput( return BatchStrOutput(
rids=["rid-0", "rid-1"], rids=["rid-0", "rid-1"],
@@ -63,6 +90,19 @@ class TestMultiTokenizerMixin(unittest.TestCase):
[{"device": 1, "host": 3}], [{"device": 1, "host": 3}],
) )
def test_get_tokenizer_worker_class_uses_default(self):
self.assertIs(get_tokenizer_worker_class(DefaultServerArgs()), TokenizerWorker)
def test_get_tokenizer_worker_class_resolves_custom_class(self):
self.assertIs(
get_tokenizer_worker_class(CustomServerArgs()),
CustomTokenizerWorker,
)
def test_get_tokenizer_worker_class_rejects_non_worker(self):
with self.assertRaisesRegex(TypeError, "TokenizerWorker"):
get_tokenizer_worker_class(InvalidServerArgs())
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()