[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:
co-authored by
tanujtiwari1998
parent
77b7698cad
commit
cfc66e05c5
@@ -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.
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user