Skip to content
Open
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
48 changes: 48 additions & 0 deletions src/tests/test_kvaware_chat_tokenization.py
Original file line number Diff line number Diff line change
Expand Up @@ -385,3 +385,51 @@ def from_pretrained(name):
ids2 = await router.tokenize_prompt(endpoints(URL_A), {"messages": MESSAGES})
assert ids2 == CHAT_IDS
assert loaded == [MODEL] # cached - no second load


# --- the --tokenizer operator override -----------------------------------------


@pytest.mark.asyncio
async def test_tokenizer_override_is_loaded_instead_of_the_served_name(monkeypatch):
"""Engines serving under an alias advertise a name that can never load;
--tokenizer supplies the real id, and the endpoint list must not even be
consulted for a name (the stub here has no model_names)."""
router = LoadAwareRouter.__new__(LoadAwareRouter)
router.tokenizer = None
router.tokenizer_name = "org/real-tokenizer-repo"
loaded = []

class FakeAuto:
@staticmethod
def from_pretrained(name):
loaded.append(name)
return ChatTokenizer()

monkeypatch.setattr(routing_logic, "AutoTokenizer", FakeAuto, raising=False)

class NamelessEndpoint:
url = URL_A # deliberately NO model_names attribute

ids = await router.tokenize_prompt([NamelessEndpoint()], {"messages": MESSAGES})
assert ids == CHAT_IDS
assert loaded == ["org/real-tokenizer-repo"]


@pytest.mark.asyncio
async def test_without_override_the_served_name_is_used(monkeypatch):
router = LoadAwareRouter.__new__(LoadAwareRouter)
router.tokenizer = None
router.tokenizer_name = None
loaded = []

class FakeAuto:
@staticmethod
def from_pretrained(name):
loaded.append(name)
return ChatTokenizer()

monkeypatch.setattr(routing_logic, "AutoTokenizer", FakeAuto, raising=False)
ids = await router.tokenize_prompt(endpoints(URL_A), {"messages": MESSAGES})
assert ids == CHAT_IDS
assert loaded == [MODEL]
1 change: 1 addition & 0 deletions src/vllm_router/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,6 +279,7 @@ def initialize_all(app: FastAPI, args):
prefill_model_labels=args.prefill_model_labels,
decode_model_labels=args.decode_model_labels,
kv_aware_threshold=args.kv_aware_threshold,
tokenizer=args.tokenizer,
loadaware_beta=args.loadaware_beta,
prefix_min_match_length=args.prefix_min_match_length,
priority_header=args.priority_header,
Expand Down
14 changes: 14 additions & 0 deletions src/vllm_router/parsers/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -449,6 +449,20 @@ def parse_args():
help="The threshold for kv-aware routing.",
)

parser.add_argument(
"--tokenizer",
type=str,
default=None,
help="Tokenizer id or local path the router loads for kv-aware/"
"load-aware token-id computation, INSTEAD of the model name the "
"engines advertise. Set this when engines serve under an alias "
"(vLLM --served-model-name) that is not a resolvable tokenizer id - "
"otherwise the router cannot tokenize locally and pays a remote "
"/tokenize round trip per routing decision. MUST be the same "
"tokenizer the engines run, or router-side token ids drift from "
"engine-side KV hashes and kv-aware routing silently degrades.",
)

parser.add_argument(
"--loadaware-beta",
type=float,
Expand Down
17 changes: 16 additions & 1 deletion src/vllm_router/routers/routing_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,7 +167,14 @@ async def _ensure_tokenizer(router, endpoints: List[EndpointInfo]):
doomed hub lookup before reaching the remote ``/tokenize`` fallback.
"""
if router.tokenizer is None:
model_name = endpoints[0].model_names[0]
# Operator override first: engines that serve under an alias
# (--served-model-name) advertise a name that is not a resolvable
# tokenizer id, which forces the remote /tokenize round trip on every
# request. --tokenizer supplies the real id (or a local path) so the
# router can tokenize in-process.
model_name = (
getattr(router, "tokenizer_name", None) or endpoints[0].model_names[0]
)
if model_name in getattr(router, "_tokenizer_load_failures", ()):
raise ValueError(
f"tokenizer load for '{model_name}' already failed; "
Expand Down Expand Up @@ -413,7 +420,11 @@ def __init__(
lmcache_worker_timeout: int = 30,
lmcache_controller_reply_port: Optional[int] = None,
lmcache_controller_heartbeat_port: Optional[int] = None,
tokenizer: Optional[str] = None,
):
#: Optional tokenizer id/path the router loads INSTEAD of the served
#: model name - must be the tokenizer the engines actually run.
self.tokenizer_name = tokenizer
self.lmcache_controller_port = lmcache_controller_port
self.lmcache_controller_reply_port = lmcache_controller_reply_port
self.lmcache_controller_heartbeat_port = lmcache_controller_heartbeat_port
Expand Down Expand Up @@ -619,6 +630,7 @@ def __init__(
lmcache_controller_reply_port: Optional[int] = None,
lmcache_controller_heartbeat_port: Optional[int] = None,
loadaware_beta: Optional[float] = None,
tokenizer: Optional[str] = None,
):
super().__init__(
lmcache_controller_port,
Expand All @@ -628,6 +640,7 @@ def __init__(
lmcache_worker_timeout=lmcache_worker_timeout,
lmcache_controller_reply_port=lmcache_controller_reply_port,
lmcache_controller_heartbeat_port=lmcache_controller_heartbeat_port,
tokenizer=tokenizer,
)
#: Weight on the load penalty, in units of "full cache hits per 100%
#: above fleet-average load".
Expand Down Expand Up @@ -1262,6 +1275,7 @@ def initialize_routing_logic(
lmcache_controller_heartbeat_port=kwargs.get(
"lmcache_controller_heartbeat_port"
),
tokenizer=kwargs.get("tokenizer"),
)
router.start_kv_manager()
elif routing_logic == RoutingLogic.LOADAWARE:
Expand All @@ -1277,6 +1291,7 @@ def initialize_routing_logic(
"lmcache_controller_heartbeat_port"
),
loadaware_beta=kwargs.get("loadaware_beta"),
tokenizer=kwargs.get("tokenizer"),
)
router.start_kv_manager()
elif routing_logic == RoutingLogic.PREFIXAWARE:
Expand Down