mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 19c5790a91 into 02f61c9c42
This commit is contained in:
commit
6cf1683259
6 changed files with 74 additions and 19 deletions
|
|
@ -5,7 +5,7 @@ from collections.abc import Iterator, Mapping, Sequence
|
|||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -71,6 +71,13 @@ class PrefetchedUsage:
|
|||
return _active_prefetched_usage.get()
|
||||
|
||||
|
||||
def _declares_tpm_limit(deployment: Mapping[str, object]) -> bool:
|
||||
"""Whether the deployment carries a ``tpm`` limit that selection has to weigh the prompt against."""
|
||||
nested: Final = (deployment.get("litellm_params"), deployment.get("model_info"))
|
||||
sources: Final = (deployment, *(source for source in nested if isinstance(source, Mapping)))
|
||||
return any(source.get("tpm") is not None for source in sources)
|
||||
|
||||
|
||||
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
||||
"""
|
||||
Updated version of TPM/RPM Logging.
|
||||
|
|
@ -421,10 +428,13 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
for idx, key in enumerate(rpm_keys):
|
||||
rpm_dict[rpm_keys[idx].split(":")[0]] = rpm_values[idx]
|
||||
|
||||
try:
|
||||
input_tokens = token_counter(messages=messages, text=input)
|
||||
except Exception:
|
||||
input_tokens = 0
|
||||
deployments: Final = cast(Sequence[Mapping[str, object]], healthy_deployments) # cast-ok: untyped dicts
|
||||
input_tokens = 0
|
||||
if any(_declares_tpm_limit(deployment) for deployment in deployments):
|
||||
try:
|
||||
input_tokens = token_counter(messages=messages, text=input)
|
||||
except Exception:
|
||||
input_tokens = 0
|
||||
verbose_router_logger.debug("input_tokens=%s", input_tokens)
|
||||
# -----------------------
|
||||
# Find lowest used model
|
||||
|
|
|
|||
|
|
@ -88,7 +88,7 @@ RULES: Final[Rules] = (
|
|||
RouteRule(Route.MESSAGES, Rollout.RUST_OPT_IN, providers=frozenset({"anthropic"})),
|
||||
RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY),
|
||||
RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY),
|
||||
RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY),
|
||||
RouteRule(Route.TOKEN_COUNTER, Rollout.RUST_OPT_IN),
|
||||
RouteRule(Route.TOKENIZER, Rollout.PYTHON_ONLY),
|
||||
RouteRule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})),
|
||||
SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.GOOGLE_KMS.value})),
|
||||
|
|
|
|||
|
|
@ -79,7 +79,7 @@ def rust_tokenizer(model: str) -> RustTokenizer | None:
|
|||
|
||||
@lru_cache(maxsize=4)
|
||||
def _counter(factory: RustTokenCounterFactory, tokenizer: RustTokenizer) -> RustTokenCounter:
|
||||
return factory.from_tokenizer(_native_tokenizer(tokenizer))
|
||||
return factory.from_tokenizer(_native_tokenizer(tokenizer), fast=True)
|
||||
|
||||
|
||||
def _native_tokenizer(tokenizer: RustTokenizer) -> NativeTokenizer:
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
from datetime import datetime, timedelta
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.router_strategy.lowest_tpm_rpm_v2 as strategy_module
|
||||
from litellm import Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
|
||||
|
|
@ -42,14 +43,11 @@ def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None:
|
|||
}
|
||||
now: Final = datetime.now()
|
||||
cache_keys: Final = tuple(
|
||||
f"{MODEL_GROUP}:tpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}"
|
||||
for offset in range(60)
|
||||
f"{MODEL_GROUP}:tpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}" for offset in range(60)
|
||||
)
|
||||
|
||||
for cache_key in cache_keys:
|
||||
router.cache.set_cache(
|
||||
key=cache_key, value=usage_by_deployment, ttl=float("inf")
|
||||
)
|
||||
router.cache.set_cache(key=cache_key, value=usage_by_deployment, ttl=float("inf"))
|
||||
|
||||
deployment: Final = router.get_available_deployment(
|
||||
model=MODEL_GROUP,
|
||||
|
|
@ -109,11 +107,58 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o
|
|||
)
|
||||
router.lowesttpm_logger_v2 = OldSignatureV2(router_cache=router.cache, routing_args={})
|
||||
|
||||
response: Final = await router.acompletion(
|
||||
model=MODEL_GROUP, messages=[{"role": "user", "content": "x"}]
|
||||
)
|
||||
response: Final = await router.acompletion(model=MODEL_GROUP, messages=[{"role": "user", "content": "x"}])
|
||||
|
||||
assert response.choices[0].message.content in {
|
||||
f"from {HIGH_USAGE_DEPLOYMENT_ID}",
|
||||
f"from {LOW_USAGE_DEPLOYMENT_ID}",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v2_counts_prompt_tokens_only_when_a_deployment_declares_a_tpm_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
counter: Final = Mock(return_value=40_000)
|
||||
monkeypatch.setattr(strategy_module, "token_counter", counter)
|
||||
router_cache: Final = DualCache()
|
||||
monkeypatch.setattr(router_cache, "async_batch_get_cache", AsyncMock(return_value=[10, 20, None, None]))
|
||||
strategy: Final = LowestTPMLoggingHandler_v2(router_cache=router_cache)
|
||||
messages: Final = [{"role": "user", "content": "a long prompt"}]
|
||||
unlimited: Final = [
|
||||
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "a"}},
|
||||
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}},
|
||||
]
|
||||
|
||||
chosen: Final = await strategy.async_get_available_deployments(
|
||||
model_group="g", healthy_deployments=unlimited, messages=messages
|
||||
)
|
||||
assert chosen["model_info"]["id"] == "a", "the lowest counter still wins without any tpm limit"
|
||||
assert counter.call_count == 0, "no deployment has a tpm limit, so the prompt is never tokenized for selection"
|
||||
|
||||
limited: Final = [
|
||||
{"model_name": "g", "litellm_params": {"model": "m", "tpm": 30_000}, "model_info": {"id": "a"}},
|
||||
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}},
|
||||
]
|
||||
chosen_limited: Final = await strategy.async_get_available_deployments(
|
||||
model_group="g", healthy_deployments=limited, messages=messages
|
||||
)
|
||||
assert counter.call_count == 1, "a declared tpm limit needs the prompt size to be enforced"
|
||||
assert chosen_limited["model_info"]["id"] == "b", "a's tpm limit cannot fit the counted prompt, so b is picked"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v2_treats_a_failed_prompt_count_as_zero_tokens(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(strategy_module, "token_counter", Mock(side_effect=ValueError("no tokenizer")))
|
||||
router_cache: Final = DualCache()
|
||||
monkeypatch.setattr(router_cache, "async_batch_get_cache", AsyncMock(return_value=[10, 20, None, None]))
|
||||
strategy: Final = LowestTPMLoggingHandler_v2(router_cache=router_cache)
|
||||
limited: Final = [
|
||||
{"model_name": "g", "litellm_params": {"model": "m", "tpm": 30_000}, "model_info": {"id": "a"}},
|
||||
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}},
|
||||
]
|
||||
|
||||
chosen: Final = await strategy.async_get_available_deployments(
|
||||
model_group="g", healthy_deployments=limited, messages=[{"role": "user", "content": "a long prompt"}]
|
||||
)
|
||||
assert chosen["model_info"]["id"] == "a", "a failed count weighs as zero tokens, so the lowest counter still wins"
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ def test_shipped_decisions(
|
|||
if route is Route.OCR or (route is Route.TRANSCRIPTION and provider == "bedrock"):
|
||||
assert catalog.rollout(context) is Rollout.RUST_REQUIRED
|
||||
assert catalog.decision(context) is Decision.RUST_REQUIRED
|
||||
elif route is Route.MESSAGES and provider == "anthropic":
|
||||
elif (route is Route.MESSAGES and provider == "anthropic") or route is Route.TOKEN_COUNTER:
|
||||
assert catalog.rollout(context) is Rollout.RUST_OPT_IN
|
||||
opted_in: Final = environment == "1" or (environment is None and process is True)
|
||||
assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if opted_in else Decision.PYTHON)
|
||||
|
|
|
|||
|
|
@ -94,7 +94,7 @@ async def test_native_count_returns_typed_count_and_reuses_one_counter(fake_toke
|
|||
assert second == first
|
||||
assert len(factory.counters) == 1
|
||||
assert factory.counters[0].bodies == [BODY, BODY]
|
||||
assert factory.counters[0].fast is False
|
||||
assert factory.counters[0].fast is True
|
||||
assert factory.counters[0].tokenizer is tokenizer_dispatch.native_anthropic()
|
||||
assert json.loads(factory.counters[0].tokenizer.json or "")["model"]["type"] == "BPE"
|
||||
|
||||
|
|
@ -113,7 +113,7 @@ async def test_tiktoken_counter_is_built_over_the_shared_encoding_once(
|
|||
assert len(factory.counters) == 1
|
||||
assert factory.counters[0].tokenizer.name == tokenizer
|
||||
assert factory.counters[0].tokenizer is tokenizer_dispatch.native_encoding(tokenizer)
|
||||
assert factory.counters[0].fast is False
|
||||
assert factory.counters[0].fast is True
|
||||
assert factory.counters[0].bodies == [BODY, BODY]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue