diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index fe6f26db6..a0d48b8f7 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -146,6 +146,7 @@ from sglang.srt.managers.multi_tokenizer_mixin import ( MultiTokenizerRouter, TokenizerWorker, get_main_process_id, + get_tokenizer_worker_class, read_from_shared_memory, write_data_for_multi_tokenizer, ) @@ -236,7 +237,8 @@ async def init_multi_tokenizer() -> ServerArgs: ) # 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.initialize_templates( tokenizer_manager=tokenizer_manager, diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 731d3e5b7..f05bba746 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -29,7 +29,7 @@ import sys import threading import zlib 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 setproctitle @@ -679,6 +679,19 @@ class TokenizerWorker(TokenizerManager): 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): """ Sometimes an asyncio function does not print exception. diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 3940bfb50..5de5fb76c 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -6670,6 +6670,11 @@ class ServerArgs: ] 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): scheme = "https" if self.ssl_certfile else "http" # When binding to all interfaces, use loopback for internal requests. diff --git a/test/registered/unit/managers/test_multi_tokenizer_mixin.py b/test/registered/unit/managers/test_multi_tokenizer_mixin.py index 5a8abfab9..a20b6d4bb 100644 --- a/test/registered/unit/managers/test_multi_tokenizer_mixin.py +++ b/test/registered/unit/managers/test_multi_tokenizer_mixin.py @@ -6,11 +6,38 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel maybe_stub_sgl_kernel() 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") +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: return BatchStrOutput( rids=["rid-0", "rid-1"], @@ -63,6 +90,19 @@ class TestMultiTokenizerMixin(unittest.TestCase): [{"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__": unittest.main()