diff --git a/src/tests/test_kvaware_router.py b/src/tests/test_kvaware_router.py index 11278920f..8506c63cd 100644 --- a/src/tests/test_kvaware_router.py +++ b/src/tests/test_kvaware_router.py @@ -22,9 +22,9 @@ def __init__(self, **kwargs): class EndpointInfo: - def __init__(self, url): + def __init__(self, url, model_name="test-model"): self.url = url - self.model_names = ["test-model"] + self.model_names = [model_name] class Tokenizer: @@ -44,7 +44,7 @@ async def test_kvaware_routes_to_longest_reported_prefix(threshold): url_a = "http://10.0.0.1:8000" url_b = "http://10.0.0.2:8000" router = KvawareRouter.__new__(KvawareRouter) - router.tokenizer = Tokenizer() + router.tokenizers = {"test-model": Tokenizer()} router.threshold = threshold router.session_key = "x-session-id" router.hash_ring = routing_logic.HashRing() @@ -79,6 +79,63 @@ async def query_manager(_msg): assert selected == url_b +@pytest.mark.asyncio +async def test_kvaware_uses_tokenizer_for_each_model(monkeypatch): + loaded_models = [] + lookup_tokens = [] + + class ModelTokenizer: + def __init__(self, token_id): + self.token_id = token_id + + def encode(self, _prompt): + return [self.token_id] + + tokenizers = { + "model-a": ModelTokenizer(1), + "model-b": ModelTokenizer(2), + } + + def load_tokenizer(model_name): + loaded_models.append(model_name) + return tokenizers[model_name] + + monkeypatch.setattr( + routing_logic, + "AutoTokenizer", + SimpleNamespace(from_pretrained=load_tokenizer), + raising=False, + ) + + router = KvawareRouter.__new__(KvawareRouter) + router.tokenizers = {} + router.threshold = 0 + router.session_key = "x-session-id" + router.hash_ring = routing_logic.HashRing() + router.instance_id_to_ip = {} + + async def query_manager(msg): + lookup_tokens.append(msg.tokens) + return SimpleNamespace(layout_info={}) + + router.query_manager = query_manager + + for model_name, url in ( + ("model-a", "http://10.0.0.1:8000"), + ("model-b", "http://10.0.0.2:8000"), + ): + await router.route_request( + [EndpointInfo(url, model_name)], + {}, + {}, + SimpleNamespace(headers={}), + {"prompt": "test"}, + ) + + assert loaded_models == ["model-a", "model-b"] + assert lookup_tokens == [[1], [2]] + + @pytest.mark.parametrize( "instance_map", [ @@ -105,7 +162,7 @@ async def test_kvaware_ignores_cached_dead_holder(instance_map): url_a = "http://10.0.0.1:8000" url_b = "http://10.0.0.2:8000" router = KvawareRouter.__new__(KvawareRouter) - router.tokenizer = Tokenizer() + router.tokenizers = {"test-model": Tokenizer()} router.threshold = 1000 router.session_key = "x-session-id" router.hash_ring = routing_logic.HashRing() @@ -136,7 +193,7 @@ async def test_kvaware_refreshes_unknown_dead_holder_and_uses_live_match(): url_a = "http://10.0.0.1:8000" url_b = "http://10.0.0.2:8000" router = KvawareRouter.__new__(KvawareRouter) - router.tokenizer = Tokenizer() + router.tokenizers = {"test-model": Tokenizer()} router.threshold = 1000 router.session_key = "x-session-id" router.hash_ring = routing_logic.HashRing() diff --git a/src/tests/test_loadaware_router.py b/src/tests/test_loadaware_router.py index 6b2596ad7..a0b4956b1 100644 --- a/src/tests/test_loadaware_router.py +++ b/src/tests/test_loadaware_router.py @@ -53,8 +53,9 @@ def __init__(self, **kwargs): class EndpointInfo: - def __init__(self, url: str): + def __init__(self, url: str, model_name: str = "test-model"): self.url = url + self.model_names = [model_name] class RequestStats: @@ -309,7 +310,7 @@ class Tokenizer: def encode(self, _prompt): return list(range(PROMPT_TOKENS)) - router.tokenizer = Tokenizer() + router.tokenizers = {"test-model": Tokenizer()} async def query_manager(_msg): return LookupRet({INST_A: (LOCAL, PROMPT_TOKENS)}) @@ -334,7 +335,7 @@ class Tokenizer: def encode(self, _prompt): return list(range(PROMPT_TOKENS)) - router.tokenizer = Tokenizer() + router.tokenizers = {"test-model": Tokenizer()} async def query_manager(_msg): return LookupRet({}) @@ -353,6 +354,46 @@ class Request: assert url == URL_B +@pytest.mark.asyncio +async def test_tokenize_prompt_uses_tokenizer_for_each_model(monkeypatch): + loaded_models = [] + + class ModelTokenizer: + def __init__(self, token_id): + self.token_id = token_id + + def encode(self, _prompt): + return [self.token_id] + + tokenizers = { + "model-a": ModelTokenizer(1), + "model-b": ModelTokenizer(2), + } + + class FakeAutoTokenizer: + @staticmethod + def from_pretrained(model_name): + loaded_models.append(model_name) + return tokenizers[model_name] + + monkeypatch.setattr( + routing_logic, "AutoTokenizer", FakeAutoTokenizer, raising=False + ) + router = make_router() + router.tokenizers = {} + + token_ids = [] + for model_name in ("model-a", "model-b", "model-a"): + token_ids.append( + await router.tokenize_prompt( + [EndpointInfo(URL_A, model_name)], {"prompt": "test"} + ) + ) + + assert loaded_models == ["model-a", "model-b"] + assert token_ids == [[1], [2], [1]] + + @pytest.mark.asyncio async def test_no_endpoints_is_a_503_not_a_crash(): from fastapi import HTTPException diff --git a/src/vllm_router/routers/routing_logic.py b/src/vllm_router/routers/routing_logic.py index 52f6998a5..81c9aa76b 100644 --- a/src/vllm_router/routers/routing_logic.py +++ b/src/vllm_router/routers/routing_logic.py @@ -326,9 +326,15 @@ def __init__( self.instance_id_to_ip = {} self.session_key = session_key self.hash_ring = HashRing() - self.tokenizer = None + self.tokenizers = {} self.threshold = kv_aware_threshold + def _get_tokenizer(self, endpoints: List[EndpointInfo]): + model_name = endpoints[0].model_names[0] + if model_name not in self.tokenizers: + self.tokenizers[model_name] = AutoTokenizer.from_pretrained(model_name) + return self.tokenizers[model_name] + def start_kv_manager(self): """ Start the kv manager @@ -389,11 +395,8 @@ async def route_request( # Local-first tokenization, fall back to remote "/tokenize" API on failure # TODO (Yuhan): Handle chat completions try: - if self.tokenizer is None: - self.tokenizer = AutoTokenizer.from_pretrained( - endpoints[0].model_names[0] - ) - token_ids = self.tokenizer.encode(request_json.get("prompt", "")) + tokenizer = self._get_tokenizer(endpoints) + token_ids = tokenizer.encode(request_json.get("prompt", "")) except Exception: # Remote /tokenize fallback (let errors bubble up to keep behavior simple) remote_url = endpoints[0].url + "/tokenize" @@ -668,11 +671,8 @@ async def tokenize_prompt( executor rather than on the event loop. """ try: - if self.tokenizer is None: - self.tokenizer = AutoTokenizer.from_pretrained( - endpoints[0].model_names[0] - ) - return self.tokenizer.encode(request_json.get("prompt", "")) + tokenizer = self._get_tokenizer(endpoints) + return tokenizer.encode(request_json.get("prompt", "")) except Exception: remote_url = endpoints[0].url + "/tokenize" headers = {"Content-Type": "application/json"}