[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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user