mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(router): add heuristic v2 complexity routing (#39276)
* feat(router): add trained heuristic complexity routing * feat(router): expose heuristic v2 classifier * style(router): format heuristic v2 predictor
This commit is contained in:
parent
082bea851e
commit
9aeeca4ce3
16 changed files with 4653 additions and 19 deletions
|
|
@ -68,6 +68,36 @@ still resolve to a deployment in `model_list`; this configuration does not creat
|
|||
- abc
|
||||
```
|
||||
|
||||
### Heuristic v2
|
||||
|
||||
Set `classifier_type: heuristic_v2` to classify with the bundled calibrated
|
||||
success-probability model instead of the hand-written weighted scorer
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: smart-router
|
||||
litellm_params:
|
||||
model: auto_router/complexity_router
|
||||
complexity_router_config:
|
||||
classifier_type: heuristic_v2
|
||||
tiers:
|
||||
SIMPLE: luna
|
||||
MEDIUM: terra
|
||||
COMPLEX: sol
|
||||
REASONING: sol-ultra
|
||||
```
|
||||
|
||||
No classifier model call or per-model training data is required. The classifier
|
||||
uses global tier quality, request-type quality, and similar-request cohorts from
|
||||
the bundled UltraFeedback artifact. It estimates success at every tier, enforces
|
||||
monotonic probabilities, and returns the first tier meeting the trained 0.75
|
||||
threshold. The existing complexity-router tier pool then selects and dispatches
|
||||
a model from that tier
|
||||
|
||||
Spend logs record `routing_decision.cause: heuristic_v2`, the detected request
|
||||
type, and all four predicted probabilities. Existing `classifier_type: heuristic`
|
||||
configurations keep the original weighted scorer unchanged
|
||||
|
||||
### Renaming the tiers
|
||||
|
||||
`tier_labels` puts your own vocabulary on the four tiers:
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -33,6 +33,11 @@ from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal
|
|||
from litellm.litellm_core_utils.prompt_templates.common_utils import request_contains_image_content
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
|
||||
from litellm.router_strategy.complexity_router.tier_predictor import (
|
||||
TierSuccessPredictor,
|
||||
resolve_tier_artifact,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
ModelResponse,
|
||||
|
|
@ -790,6 +795,7 @@ class ClassificationOutcome(NamedTuple):
|
|||
signals: tuple[str, ...]
|
||||
cause: Literal[
|
||||
"heuristic_scorer",
|
||||
"heuristic_v2",
|
||||
"reasoning_override",
|
||||
"llm_classifier",
|
||||
"heuristic_first_short_circuit",
|
||||
|
|
@ -978,6 +984,11 @@ class ComplexityRouter(CustomLogger):
|
|||
if llm_classifier_configured
|
||||
else None
|
||||
)
|
||||
self._tier_success_predictor: TierSuccessPredictor | None = (
|
||||
TierSuccessPredictor(resolve_tier_artifact(self.config.heuristic_v2_artifact))
|
||||
if self.config.classifier_type == "heuristic_v2"
|
||||
else None
|
||||
)
|
||||
|
||||
verbose_router_logger.debug("ComplexityRouter initialized for %s with tiers: %s", model_name, self.config.tiers)
|
||||
|
||||
|
|
@ -1350,6 +1361,8 @@ class ComplexityRouter(CustomLogger):
|
|||
custom tier set, and classifier_fallback otherwise decides between the heuristic scorer and
|
||||
default_model. The outcome's `cause` reports which path actually ran.
|
||||
"""
|
||||
if self.config.classifier_type == "heuristic_v2":
|
||||
return self._classify_with_heuristic_v2(prompt)
|
||||
if self.config.classifier_type == "custom":
|
||||
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
|
||||
if self.config.classifier_type == "heuristic_first" and self.config.classifier_llm_config is not None:
|
||||
|
|
@ -1359,6 +1372,24 @@ class ComplexityRouter(CustomLogger):
|
|||
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
|
||||
return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages)
|
||||
|
||||
def _classify_with_heuristic_v2(self, prompt: str) -> ClassificationOutcome:
|
||||
predictor: Final = self._tier_success_predictor
|
||||
if predictor is None:
|
||||
raise ValueError("heuristic v2 predictor is not configured")
|
||||
request_type: Final = classify_prompt(prompt)
|
||||
prediction: Final = predictor.predict(prompt, request_type)
|
||||
tier: Final = TIER_SEVERITY_ORDER[prediction.required_tier - 1]
|
||||
probability_signals: Final = tuple(
|
||||
f"tier-probability:{candidate.value.lower()}={prediction.probabilities[index]:.6f}"
|
||||
for index, candidate in enumerate(TIER_SEVERITY_ORDER, start=1)
|
||||
)
|
||||
return ClassificationOutcome(
|
||||
tier=tier,
|
||||
score=None,
|
||||
signals=(f"request-type:{request_type.value}", *probability_signals),
|
||||
cause="heuristic_v2",
|
||||
)
|
||||
|
||||
async def _classify_heuristic_first(
|
||||
self,
|
||||
prompt: str,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_seriali
|
|||
|
||||
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
|
||||
|
||||
from .tier_predictor import TrainedTierArtifact
|
||||
|
||||
|
||||
class ComplexityTier(str, Enum):
|
||||
"""Complexity tiers for routing decisions."""
|
||||
|
|
@ -625,12 +627,19 @@ class ComplexityRouterConfig(BaseModel):
|
|||
)
|
||||
|
||||
# Classifier strategy
|
||||
classifier_type: Literal["heuristic", "llm", "custom", "heuristic_first"] = Field(
|
||||
classifier_type: Literal["heuristic", "heuristic_v2", "llm", "custom", "heuristic_first"] = Field(
|
||||
default="heuristic",
|
||||
description=(
|
||||
"Classification strategy: local regex/keyword scoring, an LLM call, a custom classifier "
|
||||
"plugin, or 'heuristic_first', which scores locally and only pays for the LLM classifier "
|
||||
"when the local scorer does not confidently land a cheap tier"
|
||||
"Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, "
|
||||
"an LLM call, a custom classifier plugin, or 'heuristic_first', which scores locally and only pays "
|
||||
"for the LLM classifier when the local scorer does not confidently land a cheap tier"
|
||||
),
|
||||
)
|
||||
heuristic_v2_artifact: TrainedTierArtifact | Literal["ultrafeedback"] = Field(
|
||||
default="ultrafeedback",
|
||||
description=(
|
||||
"Success-probability artifact used by classifier_type 'heuristic_v2'. The bundled "
|
||||
"UltraFeedback artifact is selected by default; an inline trained artifact may replace it"
|
||||
),
|
||||
)
|
||||
classifier_llm_config: ClassifierLLMConfig | None = Field(
|
||||
|
|
@ -1248,10 +1257,10 @@ class ComplexityRouterConfig(BaseModel):
|
|||
)
|
||||
if duplicated:
|
||||
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
|
||||
if self.classifier_type in ("heuristic", "heuristic_first"):
|
||||
if self.classifier_type in ("heuristic", "heuristic_v2", "heuristic_first"):
|
||||
raise ValueError(
|
||||
"tier_definitions requires classifier_type 'llm' or 'custom': the heuristic scorer only "
|
||||
"produces the built-in tiers"
|
||||
"produces the four built-in tiers, as does heuristic_v2"
|
||||
)
|
||||
conflicts: Final = self._tier_definition_conflicts()
|
||||
if conflicts:
|
||||
|
|
|
|||
156
litellm/router_strategy/complexity_router/tier_predictor.py
Normal file
156
litellm/router_strategy/complexity_router/tier_predictor.py
Normal file
|
|
@ -0,0 +1,156 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from litellm.types.router import RequestType
|
||||
|
||||
|
||||
class TierGlobalStatistic(BaseModel):
|
||||
tier: int = Field(ge=1, le=4)
|
||||
successes: float = Field(ge=0.0)
|
||||
observations: float = Field(gt=0.0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _successes_do_not_exceed_observations(self) -> TierGlobalStatistic:
|
||||
if self.successes > self.observations:
|
||||
raise ValueError("successes cannot exceed observations")
|
||||
return self
|
||||
|
||||
|
||||
class TierDomainStatistic(TierGlobalStatistic):
|
||||
request_type: RequestType
|
||||
|
||||
|
||||
class TierCohortStatistic(TierGlobalStatistic):
|
||||
cohort: str = Field(min_length=1)
|
||||
|
||||
|
||||
class TierDataset(BaseModel):
|
||||
name: str = Field(min_length=1)
|
||||
url: str = Field(min_length=1)
|
||||
license: str = Field(min_length=1)
|
||||
rows: int = Field(gt=0)
|
||||
success_definition: str = Field(default="quality score meets the dataset success threshold", min_length=1)
|
||||
|
||||
|
||||
class TrainedTierArtifact(BaseModel):
|
||||
schema_version: Literal[1] = 1
|
||||
global_statistics: tuple[TierGlobalStatistic, ...]
|
||||
domain_statistics: tuple[TierDomainStatistic, ...] = ()
|
||||
cohort_statistics: tuple[TierCohortStatistic, ...] = ()
|
||||
domain_prior_mass: float = Field(default=200.0, gt=0.0)
|
||||
cohort_prior_mass: float = Field(default=20.0, gt=0.0)
|
||||
routing_threshold: float = Field(default=0.75, ge=0.0, le=1.0)
|
||||
datasets: tuple[TierDataset, ...] = ()
|
||||
success_definition: str = Field(default="quality score meets the dataset success threshold", min_length=1)
|
||||
split_method: str = Field(default="sha256(prompt): 70% train, 15% validation, 15% test", min_length=1)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _statistics_are_unique(self) -> TrainedTierArtifact:
|
||||
global_tiers: Final = tuple(stat.tier for stat in self.global_statistics)
|
||||
if frozenset(global_tiers) != frozenset((1, 2, 3, 4)) or len(global_tiers) != 4:
|
||||
raise ValueError("global statistics must contain each tier exactly once")
|
||||
domain_keys: Final = tuple((stat.request_type, stat.tier) for stat in self.domain_statistics)
|
||||
if len(domain_keys) != len(frozenset(domain_keys)):
|
||||
raise ValueError("domain statistics must contain unique request_type and tier pairs")
|
||||
cohort_keys: Final = tuple((stat.cohort, stat.tier) for stat in self.cohort_statistics)
|
||||
if len(cohort_keys) != len(frozenset(cohort_keys)):
|
||||
raise ValueError("cohort statistics must contain unique cohort and tier pairs")
|
||||
return self
|
||||
|
||||
|
||||
_CODE_PATTERN: Final = re.compile(
|
||||
r"```|\b(def|class|function|python|javascript|typescript|sql|code)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_MATH_PATTERN: Final = re.compile(
|
||||
r"\b(solve|calculate|equation|probability|theorem|proof|integral)\b|[$=]",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_MULTIPLE_CHOICE_PATTERN: Final = re.compile(r"(?:^|\s)[A-D][.)]\s")
|
||||
_TIERS: Final = (1, 2, 3, 4)
|
||||
_BUILTIN_ARTIFACTS: Final = MappingProxyType({"ultrafeedback": "ultrafeedback_tiers.json"})
|
||||
|
||||
|
||||
def resolve_tier_artifact(artifact: TrainedTierArtifact | str) -> TrainedTierArtifact:
|
||||
if isinstance(artifact, TrainedTierArtifact):
|
||||
return artifact
|
||||
filename: Final = _BUILTIN_ARTIFACTS.get(artifact)
|
||||
if filename is None:
|
||||
raise ValueError(f"unknown complexity router tier artifact: {artifact}")
|
||||
path: Final = Path(__file__).with_name("artifacts") / filename
|
||||
return TrainedTierArtifact.model_validate_json(path.read_text())
|
||||
|
||||
|
||||
def similarity_cohort(prompt: str, request_type: RequestType) -> str:
|
||||
length: Final = len(prompt)
|
||||
length_bucket: Final = (
|
||||
"short" if length < 200 else "medium" if length < 800 else "long" if length < 2000 else "very_long"
|
||||
)
|
||||
code: Final = int(bool(_CODE_PATTERN.search(prompt)))
|
||||
math: Final = int(bool(_MATH_PATTERN.search(prompt)))
|
||||
multiple_choice: Final = int(bool(_MULTIPLE_CHOICE_PATTERN.search(prompt)))
|
||||
non_ascii: Final = int(sum(ord(character) > 127 for character in prompt) / max(1, length) > 0.1)
|
||||
return f"{request_type.value}|{length_bucket}|code={code}|math={math}|mc={multiple_choice}|intl={non_ascii}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TierPrediction:
|
||||
probabilities: Mapping[int, float]
|
||||
required_tier: int
|
||||
|
||||
|
||||
class TierSuccessPredictor:
|
||||
def __init__(self, artifact: TrainedTierArtifact) -> None:
|
||||
self._artifact = artifact
|
||||
self._global: Mapping[int, TierGlobalStatistic] = MappingProxyType(
|
||||
{stat.tier: stat for stat in artifact.global_statistics}
|
||||
)
|
||||
self._domain: Mapping[tuple[RequestType, int], TierDomainStatistic] = MappingProxyType(
|
||||
{(stat.request_type, stat.tier): stat for stat in artifact.domain_statistics}
|
||||
)
|
||||
self._cohort: Mapping[tuple[str, int], TierCohortStatistic] = MappingProxyType(
|
||||
{(stat.cohort, stat.tier): stat for stat in artifact.cohort_statistics}
|
||||
)
|
||||
|
||||
@property
|
||||
def routing_threshold(self) -> float:
|
||||
return self._artifact.routing_threshold
|
||||
|
||||
def predict(self, prompt: str, request_type: RequestType) -> TierPrediction:
|
||||
cohort: Final = similarity_cohort(prompt, request_type)
|
||||
raw: Final = tuple(self._probability(tier, request_type, cohort) for tier in _TIERS)
|
||||
monotonic: Final = tuple(max(raw[:index]) for index in range(1, len(raw) + 1))
|
||||
probabilities: Final[Mapping[int, float]] = MappingProxyType(
|
||||
{int(tier): probability for tier, probability in zip(_TIERS, monotonic)}
|
||||
)
|
||||
required_tier: Final = next(
|
||||
(tier for tier in _TIERS if probabilities[tier] >= self._artifact.routing_threshold),
|
||||
4,
|
||||
)
|
||||
return TierPrediction(probabilities=probabilities, required_tier=required_tier)
|
||||
|
||||
def _probability(self, tier: int, request_type: RequestType, cohort: str) -> float:
|
||||
global_stat: Final = self._global[tier]
|
||||
global_mean: Final = (global_stat.successes + 1.0) / (global_stat.observations + 2.0)
|
||||
domain_stat: Final = self._domain.get((request_type, tier))
|
||||
domain_mean: Final = self._posterior_mean(domain_stat, self._artifact.domain_prior_mass, global_mean)
|
||||
cohort_stat: Final = self._cohort.get((cohort, tier))
|
||||
return self._posterior_mean(cohort_stat, self._artifact.cohort_prior_mass, domain_mean)
|
||||
|
||||
@staticmethod
|
||||
def _posterior_mean(
|
||||
statistic: TierGlobalStatistic | None,
|
||||
prior_mass: float,
|
||||
prior_mean: float,
|
||||
) -> float:
|
||||
if statistic is None:
|
||||
return prior_mean
|
||||
return (statistic.successes + prior_mass * prior_mean) / (statistic.observations + prior_mass)
|
||||
|
|
@ -2839,6 +2839,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict):
|
|||
|
||||
RoutingDecisionCause = Literal[
|
||||
"heuristic_scorer",
|
||||
"heuristic_v2",
|
||||
# The scorer found 2+ reasoning markers and forced REASONING regardless of score.
|
||||
# A distinct cause rather than a marker inside `signals`, because it is the fact
|
||||
# that tells a reader the score did NOT choose the tier; encoding it as free text
|
||||
|
|
|
|||
|
|
@ -278,7 +278,10 @@ bindings = "pyo3"
|
|||
features = ["extension-module"]
|
||||
profile = "release"
|
||||
editable-profile = "dev"
|
||||
include = ["litellm/proxy/_experimental/out/**"]
|
||||
include = [
|
||||
"litellm/proxy/_experimental/out/**",
|
||||
"litellm/router_strategy/complexity_router/artifacts/*.json",
|
||||
]
|
||||
exclude = [
|
||||
"litellm/proxy/enterprise",
|
||||
"litellm/proxy/enterprise/**",
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_router_logger
|
||||
|
|
@ -34,10 +33,14 @@ from litellm.router_strategy.complexity_router.config import (
|
|||
DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE,
|
||||
DEFAULT_COMPLEXITY_CONFIG,
|
||||
DEFAULT_TECHNICAL_KEYWORDS,
|
||||
ClassificationRubric,
|
||||
ClassifierLLMConfig,
|
||||
ComplexityRouterConfig,
|
||||
ComplexityTier,
|
||||
ClassificationRubric,
|
||||
)
|
||||
from litellm.router_strategy.complexity_router.tier_predictor import (
|
||||
TierGlobalStatistic,
|
||||
TrainedTierArtifact,
|
||||
)
|
||||
from litellm.types.router import (
|
||||
Deployment,
|
||||
|
|
@ -46,6 +49,16 @@ from litellm.types.router import (
|
|||
)
|
||||
|
||||
|
||||
def _heuristic_v2_artifact() -> TrainedTierArtifact:
|
||||
return TrainedTierArtifact(
|
||||
global_statistics=tuple(
|
||||
TierGlobalStatistic(tier=tier, successes=successes, observations=100)
|
||||
for tier, successes in enumerate((10, 20, 90, 99), start=1)
|
||||
),
|
||||
routing_threshold=0.8,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_router_instance():
|
||||
"""Create a mock LiteLLM Router instance."""
|
||||
|
|
@ -1696,6 +1709,59 @@ class TestLLMClassifier:
|
|||
assert outcome.cause == "heuristic_scorer"
|
||||
assert outcome.score is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heuristic_v2_routes_directly_to_predicted_builtin_tier(self, mock_router_instance):
|
||||
router = ComplexityRouter(
|
||||
model_name="tier-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"classifier_type": "heuristic_v2",
|
||||
"heuristic_v2_artifact": _heuristic_v2_artifact(),
|
||||
"tiers": {
|
||||
"SIMPLE": "simple-model",
|
||||
"MEDIUM": "medium-model",
|
||||
"COMPLEX": "complex-model",
|
||||
"REASONING": "reasoning-model",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
response = await router.async_pre_routing_hook(
|
||||
model="tier-router",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "Handle this new request"}],
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.model == "complex-model"
|
||||
assert response.routing_decision["tier"] == "COMPLEX"
|
||||
assert response.routing_decision["cause"] == "heuristic_v2"
|
||||
assert response.routing_decision["signals"] == [
|
||||
"request-type:general",
|
||||
"tier-probability:simple=0.107843",
|
||||
"tier-probability:medium=0.205882",
|
||||
"tier-probability:complex=0.892157",
|
||||
"tier-probability:reasoning=0.980392",
|
||||
]
|
||||
|
||||
def test_heuristic_v2_needs_no_classifier_model(self):
|
||||
config = ComplexityRouterConfig(classifier_type="heuristic_v2")
|
||||
|
||||
assert config.classifier_llm_config is None
|
||||
assert config.heuristic_v2_artifact == "ultrafeedback"
|
||||
|
||||
def test_heuristic_v2_rejects_custom_tier_definitions(self):
|
||||
with pytest.raises(ValidationError, match="as does heuristic_v2"):
|
||||
ComplexityRouterConfig(
|
||||
classifier_type="heuristic_v2",
|
||||
tier_definitions=(
|
||||
{"name": "low", "description": "easy work"},
|
||||
{"name": "high", "description": "hard work"},
|
||||
),
|
||||
tiers={"low": "cheap", "high": "expensive"},
|
||||
fallback_tier="high",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aclassify_llm_success_routes_by_llm_verdict(self, llm_complexity_router, mock_router_instance):
|
||||
"""A well-formed structured LLM response should decide the tier directly.
|
||||
|
|
|
|||
|
|
@ -0,0 +1,91 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.router_strategy.complexity_router.tier_predictor import (
|
||||
TierCohortStatistic,
|
||||
TierDomainStatistic,
|
||||
TierGlobalStatistic,
|
||||
TierSuccessPredictor,
|
||||
TrainedTierArtifact,
|
||||
resolve_tier_artifact,
|
||||
similarity_cohort,
|
||||
)
|
||||
from litellm.types.router import RequestType
|
||||
|
||||
|
||||
def _artifact(
|
||||
global_successes: tuple[float, float, float, float] = (4.0, 5.0, 6.0, 7.0),
|
||||
threshold: float = 0.75,
|
||||
domain_statistics: tuple[TierDomainStatistic, ...] = (),
|
||||
cohort_statistics: tuple[TierCohortStatistic, ...] = (),
|
||||
) -> TrainedTierArtifact:
|
||||
return TrainedTierArtifact(
|
||||
global_statistics=tuple(
|
||||
TierGlobalStatistic(tier=tier, successes=successes, observations=10.0)
|
||||
for tier, successes in enumerate(global_successes, start=1)
|
||||
),
|
||||
domain_statistics=domain_statistics,
|
||||
cohort_statistics=cohort_statistics,
|
||||
domain_prior_mass=10.0,
|
||||
cohort_prior_mass=10.0,
|
||||
routing_threshold=threshold,
|
||||
)
|
||||
|
||||
|
||||
def test_predictions_are_monotonic_across_tiers() -> None:
|
||||
predictor: Final = TierSuccessPredictor(_artifact(global_successes=(9.0, 2.0, 7.0, 6.0)))
|
||||
|
||||
prediction: Final = predictor.predict("hello", RequestType.GENERAL)
|
||||
|
||||
probabilities: Final = tuple(prediction.probabilities.values())
|
||||
assert probabilities == tuple(sorted(probabilities))
|
||||
|
||||
|
||||
def test_domain_and_cohort_statistics_back_off_hierarchically() -> None:
|
||||
matching_cohort: Final = similarity_cohort("hello", RequestType.GENERAL)
|
||||
artifact: Final = _artifact(
|
||||
global_successes=(1.0, 5.0, 6.0, 7.0),
|
||||
domain_statistics=(
|
||||
TierDomainStatistic(
|
||||
tier=1,
|
||||
request_type=RequestType.GENERAL,
|
||||
successes=10.0,
|
||||
observations=10.0,
|
||||
),
|
||||
),
|
||||
cohort_statistics=(
|
||||
TierCohortStatistic(
|
||||
tier=1,
|
||||
cohort=matching_cohort,
|
||||
successes=0.0,
|
||||
observations=10.0,
|
||||
),
|
||||
),
|
||||
)
|
||||
predictor: Final = TierSuccessPredictor(artifact)
|
||||
|
||||
cohort_probability: Final = predictor.predict("hello", RequestType.GENERAL).probabilities[1]
|
||||
domain_probability: Final = predictor.predict("hello " * 100, RequestType.GENERAL).probabilities[1]
|
||||
global_probability: Final = predictor.predict("hello", RequestType.WRITING).probabilities[1]
|
||||
|
||||
assert cohort_probability == pytest.approx(7.0 / 24.0)
|
||||
assert domain_probability == pytest.approx(7.0 / 12.0)
|
||||
assert global_probability == pytest.approx(1.0 / 6.0)
|
||||
|
||||
|
||||
def test_selects_first_tier_above_probability_threshold() -> None:
|
||||
predictor: Final = TierSuccessPredictor(_artifact(global_successes=(4.0, 6.0, 8.0, 9.0), threshold=0.7))
|
||||
|
||||
prediction: Final = predictor.predict("hello", RequestType.GENERAL)
|
||||
|
||||
assert prediction.required_tier == 3
|
||||
|
||||
|
||||
def test_builtin_ultrafeedback_artifact_is_loadable() -> None:
|
||||
artifact: Final = resolve_tier_artifact("ultrafeedback")
|
||||
|
||||
assert artifact.routing_threshold == 0.75
|
||||
assert artifact.domain_prior_mass == 200.0
|
||||
assert artifact.cohort_prior_mass == 20.0
|
||||
assert artifact.datasets[0].license == "MIT"
|
||||
|
|
@ -43,6 +43,10 @@ const DEFAULT_SCORING_EXPLANATION =
|
|||
"The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical " +
|
||||
"terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the tier:";
|
||||
|
||||
const HEURISTIC_V2_EXPLANATION =
|
||||
"The router estimates success probability for all four tiers with the bundled calibrated model, then selects " +
|
||||
"the first tier that meets its trained threshold. It runs locally with no classifier API call.";
|
||||
|
||||
const CLASSIFIER_TIMEOUT_ID = "classifier-timeout-ms";
|
||||
const CLASSIFIER_CONTEXT_WINDOW_SIZE_ID = "classifier-context-window-size";
|
||||
const CLASSIFIER_CONTEXT_BUDGET_CHARS_ID = "classifier-context-budget-chars";
|
||||
|
|
@ -62,6 +66,7 @@ const CUSTOM_PROMPT_WITH_DEFAULT_MODEL_FALLBACK =
|
|||
* at all, so the panel must not keep implying a score is involved on either router.
|
||||
*/
|
||||
const scoringExplanation = (value: ComplexityRouterConfigValue): string => {
|
||||
if (value.classifier_type === "heuristic_v2") return HEURISTIC_V2_EXPLANATION;
|
||||
const usesCustomPrompt =
|
||||
usesLlmClassifier(value.classifier_type) && Boolean(value.classifier_llm_config?.system_prompt?.trim());
|
||||
if (!usesCustomPrompt) return DEFAULT_SCORING_EXPLANATION;
|
||||
|
|
@ -179,6 +184,17 @@ const ClassifierTypeRadios: React.FC<{
|
|||
</span>
|
||||
</Label>
|
||||
</SimpleTooltip>
|
||||
<SimpleTooltip content={scorerLockedReason}>
|
||||
<Label className="items-start font-normal leading-normal has-data-disabled:cursor-not-allowed has-data-disabled:opacity-50">
|
||||
<RadioGroupItem value="heuristic_v2" className="mt-0.5" disabled={scorerLocked} />
|
||||
<span>
|
||||
<strong className="font-semibold">Heuristic v2</strong>{" "}
|
||||
<span className="text-muted-foreground">
|
||||
uses bundled calibrated four-tier probabilities with no API call
|
||||
</span>
|
||||
</span>
|
||||
</Label>
|
||||
</SimpleTooltip>
|
||||
<Label className="items-start font-normal leading-normal">
|
||||
<RadioGroupItem value="llm" className="mt-0.5" />
|
||||
<span>
|
||||
|
|
|
|||
|
|
@ -127,6 +127,31 @@ describe("ComplexityRouterConfig", () => {
|
|||
expect(onChange).toHaveBeenCalledWith(expectedValue);
|
||||
});
|
||||
|
||||
it("selects heuristic v2 without requiring a classifier model or showing weighted scoring", () => {
|
||||
const onChange = vi.fn();
|
||||
const { rerender } = renderWithProviders(
|
||||
<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={onChange} />,
|
||||
);
|
||||
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
fireEvent.click(screen.getByText("Heuristic v2"));
|
||||
|
||||
expect(onChange).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
classifier_type: "heuristic_v2",
|
||||
classifier_llm_config: undefined,
|
||||
}),
|
||||
);
|
||||
|
||||
const heuristicV2Value: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "heuristic_v2" };
|
||||
rerender(<ComplexityRouterConfig modelInfo={mockModelInfo} value={heuristicV2Value} onChange={onChange} />);
|
||||
|
||||
expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("Advanced scoring")).not.toBeInTheDocument();
|
||||
expect(screen.getByText(/estimates success probability for all four tiers/)).toBeInTheDocument();
|
||||
expect(screen.queryByText(/Score < 0.15/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show classifier fields and use the configured values when classifier_type is llm", () => {
|
||||
const llmValue: ComplexityRouterConfigValue = {
|
||||
...defaultValue,
|
||||
|
|
|
|||
|
|
@ -128,7 +128,7 @@ export interface ClassifierLLMConfig {
|
|||
system_prompt?: string;
|
||||
}
|
||||
|
||||
export type ClassifierType = "heuristic" | "llm" | "heuristic_first";
|
||||
export type ClassifierType = "heuristic" | "heuristic_v2" | "llm" | "heuristic_first";
|
||||
|
||||
/**
|
||||
* Whether this router can call classifier_llm_config.model. Mirrors the backend's
|
||||
|
|
@ -161,6 +161,7 @@ export const heuristicScoringRoleFor = (
|
|||
classifierType: ClassifierType,
|
||||
classifierFallback: ClassifierFallback | undefined,
|
||||
): HeuristicScoringRole => {
|
||||
if (classifierType === "heuristic_v2") return "never";
|
||||
if (classifierType === "heuristic" || classifierType === "heuristic_first") return "decides";
|
||||
return (classifierFallback ?? DEFAULT_CLASSIFIER_FALLBACK) === "heuristic" ? "fallback_only" : "never";
|
||||
};
|
||||
|
|
@ -188,13 +189,19 @@ const builtInTierInfo = (rowId: string): { label: string; description: string; e
|
|||
return builtIn ? TIER_DESCRIPTIONS[builtIn] : undefined;
|
||||
};
|
||||
|
||||
const tierConfigIntroText = (value: ComplexityRouterConfigValue): string => {
|
||||
if (value.classifier_type === "heuristic_v2") {
|
||||
return "The complexity router classifies each request with a calibrated local four-tier model (no API calls). Configure which model(s) handle each tier.";
|
||||
}
|
||||
if (heuristicScoringRole(value) === "never") {
|
||||
return "The complexity router classifies each request with your classifier model and routes it to that tier. Configure which model(s) handle each tier.";
|
||||
}
|
||||
return "The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls, <1ms latency). Configure which model(s) handle each tier.";
|
||||
};
|
||||
|
||||
const TierConfigIntro: React.FC<{ value: ComplexityRouterConfigValue }> = ({ value }) => (
|
||||
<>
|
||||
<span className="block mb-6 text-muted-foreground">
|
||||
{heuristicScoringRole(value) === "never"
|
||||
? "The complexity router classifies each request with your classifier model and routes it to that tier. Configure which model(s) handle each tier."
|
||||
: "The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls, <1ms latency). Configure which model(s) handle each tier."}
|
||||
</span>
|
||||
<span className="block mb-6 text-muted-foreground">{tierConfigIntroText(value)}</span>
|
||||
|
||||
<span className="block mb-4 text-xs text-muted-foreground">
|
||||
{restrictedBy(value, "displayNames")?.reason ??
|
||||
|
|
|
|||
|
|
@ -165,6 +165,7 @@ describe("ClassificationMethodConfig scorer gating", () => {
|
|||
|
||||
it.each([
|
||||
["heuristic decides the tier", "heuristic" as ClassifierType, undefined, true],
|
||||
["heuristic v2 decides without the weighted scorer", "heuristic_v2" as ClassifierType, undefined, false],
|
||||
["an LLM classifier falls back to the heuristic", "llm" as ClassifierType, "heuristic" as ClassifierFallback, true],
|
||||
[
|
||||
"an LLM classifier falls back to the default model",
|
||||
|
|
|
|||
|
|
@ -111,6 +111,21 @@ describe("buildComplexityRouterConfig", () => {
|
|||
expect(config.classifier_llm_config).toBeUndefined();
|
||||
});
|
||||
|
||||
it("emits heuristic_v2 without classifier-only fields", () => {
|
||||
const trainedParams: BuildComplexityRouterConfigParams = {
|
||||
...baseParams,
|
||||
classifierType: "heuristic_v2",
|
||||
classifierLlmConfig: { model: "gpt-4o-mini", timeout_ms: 3000 },
|
||||
classifierContextWindowSize: 5,
|
||||
classifierFallback: "heuristic",
|
||||
};
|
||||
const config = buildComplexityRouterConfig(trainedParams);
|
||||
expect(config.classifier_type).toBe("heuristic_v2");
|
||||
expect(config.classifier_llm_config).toBeUndefined();
|
||||
expect(config.classifier_context_window_size).toBeUndefined();
|
||||
expect(config.classifier_fallback).toBeUndefined();
|
||||
});
|
||||
|
||||
it("includes classifier_context_window_size and classifier_context_budget_chars only when classifier_type is llm", () => {
|
||||
const params: BuildComplexityRouterConfigParams = {
|
||||
...baseParams,
|
||||
|
|
@ -713,6 +728,10 @@ describe("getClassifierModelError", () => {
|
|||
expect(getClassifierModelError({ classifier_type: "heuristic" })).toBeNull();
|
||||
});
|
||||
|
||||
it("stays quiet for a heuristic v2 router, which runs locally", () => {
|
||||
expect(getClassifierModelError({ classifier_type: "heuristic_v2" })).toBeNull();
|
||||
});
|
||||
|
||||
it("blocks an LLM classifier with no model, which the router cannot start without", () => {
|
||||
expect(getClassifierModelError({ classifier_type: "llm" })).toBe(
|
||||
"Please select a classifier model, or switch back to Heuristic",
|
||||
|
|
@ -791,7 +810,7 @@ describe("heuristic_first", () => {
|
|||
});
|
||||
|
||||
it("omits heuristic_first_max_tier on every other classifier type, which the backend rejects it on", () => {
|
||||
for (const classifierType of ["heuristic", "llm"] as const) {
|
||||
for (const classifierType of ["heuristic", "heuristic_v2", "llm"] as const) {
|
||||
const config = buildComplexityRouterConfig({ ...heuristicFirstParams, classifierType });
|
||||
expect(config.heuristic_first_max_tier).toBeUndefined();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -84,6 +84,7 @@ function describeReasoningOverride(tierLabel: string | undefined, floor: number
|
|||
|
||||
const CONSTANT_CAUSE_LABELS: Record<string, string> = {
|
||||
heuristic_scorer: "Heuristic scorer",
|
||||
heuristic_v2: "Heuristic v2",
|
||||
heuristic_first_short_circuit: "Heuristic scorer, classifier skipped",
|
||||
classifier_plugin: "Custom classifier plugin",
|
||||
semantic_keyword_match: "Semantic keyword match",
|
||||
|
|
|
|||
115
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
115
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -34575,11 +34575,11 @@ export interface components {
|
|||
classifier_plugin_timeout_ms: number;
|
||||
/**
|
||||
* Classifier Type
|
||||
* @description Classification strategy: local regex/keyword scoring, an LLM call, a custom classifier plugin, or 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier
|
||||
* @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, an LLM call, a custom classifier plugin, or 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier
|
||||
* @default heuristic
|
||||
* @enum {string}
|
||||
*/
|
||||
classifier_type: "heuristic" | "llm" | "custom" | "heuristic_first";
|
||||
classifier_type: "heuristic" | "heuristic_v2" | "llm" | "custom" | "heuristic_first";
|
||||
/**
|
||||
* Code Keywords
|
||||
* @description Keywords indicating code-related content
|
||||
|
|
@ -34640,6 +34640,12 @@ export interface components {
|
|||
* @description The highest tier the local scorer may decide on its own; required when classifier_type is 'heuristic_first' and rejected otherwise. A request whose heuristic tier is at or below this one skips the LLM classifier and routes straight to that heuristic tier, so the classifier call is only paid for on traffic the scorer could not place cheaply. The scorer must also have produced at least one signal: a prompt where no dimension fired scores 0.0 and would otherwise land SIMPLE by default rather than by evidence, which is how a chained router would silently send unclassified traffic to the cheapest model. Names a built-in tier, and may not name the highest one, since that would make the LLM classifier unreachable.
|
||||
*/
|
||||
heuristic_first_max_tier?: string | null;
|
||||
/**
|
||||
* Heuristic V2 Artifact
|
||||
* @description Success-probability artifact used by classifier_type 'heuristic_v2'. The bundled UltraFeedback artifact is selected by default; an inline trained artifact may replace it
|
||||
* @default ultrafeedback
|
||||
*/
|
||||
heuristic_v2_artifact: components["schemas"]["TrainedTierArtifact"] | "ultrafeedback";
|
||||
/**
|
||||
* Housekeeping Patterns
|
||||
* @description Additional case-sensitive literal sentinels that mark a request as client housekeeping, on top of the built-in conversation-title ones. For clients whose wording the built-ins don't cover, or after a client release changes its strings.
|
||||
|
|
@ -34778,6 +34784,12 @@ export interface components {
|
|||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
/**
|
||||
* RequestType
|
||||
* @description Fixed v0 taxonomy. User-extensible types come in v1.
|
||||
* @enum {string}
|
||||
*/
|
||||
RequestType: "code_generation" | "code_understanding" | "technical_design" | "analytical_reasoning" | "writing" | "factual_lookup" | "general";
|
||||
/** ResetSpendRequest */
|
||||
ResetSpendRequest: {
|
||||
/** Reset To */
|
||||
|
|
@ -35855,7 +35867,7 @@ export interface components {
|
|||
* Cause
|
||||
* @enum {string}
|
||||
*/
|
||||
cause?: "heuristic_scorer" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
|
||||
cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
|
||||
/** Classifier Cost */
|
||||
classifier_cost?: number;
|
||||
/** Classifier Model */
|
||||
|
|
@ -36738,6 +36750,33 @@ export interface components {
|
|||
[key: string]: unknown;
|
||||
};
|
||||
};
|
||||
/** TierCohortStatistic */
|
||||
TierCohortStatistic: {
|
||||
/** Cohort */
|
||||
cohort: string;
|
||||
/** Observations */
|
||||
observations: number;
|
||||
/** Successes */
|
||||
successes: number;
|
||||
/** Tier */
|
||||
tier: number;
|
||||
};
|
||||
/** TierDataset */
|
||||
TierDataset: {
|
||||
/** License */
|
||||
license: string;
|
||||
/** Name */
|
||||
name: string;
|
||||
/** Rows */
|
||||
rows: number;
|
||||
/**
|
||||
* Success Definition
|
||||
* @default quality score meets the dataset success threshold
|
||||
*/
|
||||
success_definition: string;
|
||||
/** Url */
|
||||
url: string;
|
||||
};
|
||||
/**
|
||||
* TierDefinition
|
||||
* @description An operator-defined tier: the name the LLM classifier must return and its rubric description.
|
||||
|
|
@ -36754,6 +36793,25 @@ export interface components {
|
|||
*/
|
||||
name: string;
|
||||
};
|
||||
/** TierDomainStatistic */
|
||||
TierDomainStatistic: {
|
||||
/** Observations */
|
||||
observations: number;
|
||||
request_type: components["schemas"]["RequestType"];
|
||||
/** Successes */
|
||||
successes: number;
|
||||
/** Tier */
|
||||
tier: number;
|
||||
};
|
||||
/** TierGlobalStatistic */
|
||||
TierGlobalStatistic: {
|
||||
/** Observations */
|
||||
observations: number;
|
||||
/** Successes */
|
||||
successes: number;
|
||||
/** Tier */
|
||||
tier: number;
|
||||
};
|
||||
/**
|
||||
* TokenCountDetailsResponse
|
||||
* @description Response structure for token count details with modality breakdown.
|
||||
|
|
@ -37023,6 +37081,57 @@ export interface components {
|
|||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
/** TrainedTierArtifact */
|
||||
TrainedTierArtifact: {
|
||||
/**
|
||||
* Cohort Prior Mass
|
||||
* @default 20
|
||||
*/
|
||||
cohort_prior_mass: number;
|
||||
/**
|
||||
* Cohort Statistics
|
||||
* @default []
|
||||
*/
|
||||
cohort_statistics: components["schemas"]["TierCohortStatistic"][];
|
||||
/**
|
||||
* Datasets
|
||||
* @default []
|
||||
*/
|
||||
datasets: components["schemas"]["TierDataset"][];
|
||||
/**
|
||||
* Domain Prior Mass
|
||||
* @default 200
|
||||
*/
|
||||
domain_prior_mass: number;
|
||||
/**
|
||||
* Domain Statistics
|
||||
* @default []
|
||||
*/
|
||||
domain_statistics: components["schemas"]["TierDomainStatistic"][];
|
||||
/** Global Statistics */
|
||||
global_statistics: components["schemas"]["TierGlobalStatistic"][];
|
||||
/**
|
||||
* Routing Threshold
|
||||
* @default 0.75
|
||||
*/
|
||||
routing_threshold: number;
|
||||
/**
|
||||
* Schema Version
|
||||
* @default 1
|
||||
* @constant
|
||||
*/
|
||||
schema_version: 1;
|
||||
/**
|
||||
* Split Method
|
||||
* @default sha256(prompt): 70% train, 15% validation, 15% test
|
||||
*/
|
||||
split_method: string;
|
||||
/**
|
||||
* Success Definition
|
||||
* @default quality score meets the dataset success threshold
|
||||
*/
|
||||
success_definition: string;
|
||||
};
|
||||
/** TransformRequestBody */
|
||||
TransformRequestBody: {
|
||||
call_type: components["schemas"]["CallTypes"];
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue