Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 62 additions & 5 deletions src/tests/test_kvaware_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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()
Expand Down Expand Up @@ -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",
[
Expand All @@ -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()
Expand Down Expand Up @@ -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()
Expand Down
47 changes: 44 additions & 3 deletions src/tests/test_loadaware_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)})
Expand All @@ -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({})
Expand All @@ -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
Expand Down
22 changes: 11 additions & 11 deletions src/vllm_router/routers/routing_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Comment thread
dsxyy marked this conversation as resolved.

def start_kv_manager(self):
"""
Start the kv manager
Expand Down Expand Up @@ -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", ""))
Comment thread
dsxyy marked this conversation as resolved.
except Exception:
# Remote /tokenize fallback (let errors bubble up to keep behavior simple)
remote_url = endpoints[0].url + "/tokenize"
Expand Down Expand Up @@ -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)
Comment thread
dsxyy marked this conversation as resolved.
return tokenizer.encode(request_json.get("prompt", ""))
Comment thread
dsxyy marked this conversation as resolved.
except Exception:
remote_url = endpoints[0].url + "/tokenize"
headers = {"Content-Type": "application/json"}
Expand Down
Loading