From 8f32662125da09c5320f0c5843d35b86ecb28a36 Mon Sep 17 00:00:00 2001 From: Tin Date: Tue, 1 Sep 2026 21:59:00 -0700 Subject: [PATCH] feat(router): wire trained tiers into complexity routing --- .../complexity_router/README.md | 30 +++++ .../complexity_router/complexity_router.py | 31 +++++ .../complexity_router/config.py | 26 ++-- litellm/types/utils.py | 1 + .../router_strategy/test_complexity_router.py | 68 ++++++++++- .../add_model/ClassificationMethodConfig.tsx | 16 +++ .../add_model/ComplexityRouterConfig.test.tsx | 25 ++++ .../add_model/ComplexityRouterConfig.tsx | 19 ++- .../add_model/HeuristicScoringConfig.test.tsx | 6 + .../build_complexity_router_config.test.ts | 21 +++- .../LogDetailsDrawer/RoutingDecisionCard.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 115 +++++++++++++++++- 12 files changed, 340 insertions(+), 19 deletions(-) diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index bc8df67cc28..bac43636702 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -68,6 +68,36 @@ still resolve to a deployment in `model_list`; this configuration does not creat - abc ``` +### Trained four-tier heuristic + +Set `classifier_type: trained_heuristic` 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: trained_heuristic + 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: trained_heuristic`, 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: diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 577cee0920d..632d6ca8c3b 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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.adaptive_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", + "trained_heuristic", "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.trained_heuristic_artifact)) + if self.config.classifier_type == "trained_heuristic" + 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 == "trained_heuristic": + return self._classify_with_trained_heuristic(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_trained_heuristic(self, prompt: str) -> ClassificationOutcome: + predictor: Final = self._tier_success_predictor + if predictor is None: + raise ValueError("trained heuristic 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="trained_heuristic", + ) + async def _classify_heuristic_first( self, prompt: str, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 70aeecb31c6..6e5082584bc 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -12,7 +12,12 @@ from typing import Annotated, Final, Literal from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_serializer, field_validator, model_validator -from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin +from litellm.types.router import ( + AdaptiveRouterTierArtifact, + AdaptiveRouterWeights, + ClassifierPlugin, + RoutingPlugin, +) class ComplexityTier(str, Enum): @@ -625,12 +630,19 @@ class ComplexityRouterConfig(BaseModel): ) # Classifier strategy - classifier_type: Literal["heuristic", "llm", "custom", "heuristic_first"] = Field( + classifier_type: Literal["heuristic", "trained_heuristic", "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" + ), + ) + trained_heuristic_artifact: AdaptiveRouterTierArtifact | Literal["ultrafeedback"] = Field( + default="ultrafeedback", + description=( + "Success-probability artifact used by classifier_type 'trained_heuristic'. The bundled " + "UltraFeedback artifact is selected by default; an inline trained artifact may replace it" ), ) classifier_llm_config: ClassifierLLMConfig | None = Field( @@ -1248,10 +1260,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", "trained_heuristic", "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 trained_heuristic" ) conflicts: Final = self._tier_definition_conflicts() if conflicts: diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 5783a39b30c..8419caa5e08 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2839,6 +2839,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict): RoutingDecisionCause = Literal[ "heuristic_scorer", + "trained_heuristic", # 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 diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 1ec8be88c9b..672b9f0142a 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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,18 +33,30 @@ 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.types.router import ( + AdaptiveRouterTierArtifact, + AdaptiveRouterTierGlobalStatistic, Deployment, LiteLLM_Params, TaggedPreRoutingStrategy, ) +def _trained_heuristic_artifact() -> AdaptiveRouterTierArtifact: + return AdaptiveRouterTierArtifact( + global_statistics=tuple( + AdaptiveRouterTierGlobalStatistic(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 +1707,59 @@ class TestLLMClassifier: assert outcome.cause == "heuristic_scorer" assert outcome.score is not None + @pytest.mark.asyncio + async def test_trained_heuristic_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": "trained_heuristic", + "trained_heuristic_artifact": _trained_heuristic_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"] == "trained_heuristic" + 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_trained_heuristic_needs_no_classifier_model(self): + config = ComplexityRouterConfig(classifier_type="trained_heuristic") + + assert config.classifier_llm_config is None + assert config.trained_heuristic_artifact == "ultrafeedback" + + def test_trained_heuristic_rejects_custom_tier_definitions(self): + with pytest.raises(ValidationError, match="as does trained_heuristic"): + ComplexityRouterConfig( + classifier_type="trained_heuristic", + 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. diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index 96c93306611..592a3f79abe 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -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 TRAINED_HEURISTIC_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 === "trained_heuristic") return TRAINED_HEURISTIC_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<{ + + +