feat(router): wire trained tiers into complexity routing

This commit is contained in:
Tin 2026-09-01 21:59:00 -07:00
parent f44e111a12
commit 8f32662125
12 changed files with 340 additions and 19 deletions

View file

@ -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:

View file

@ -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,

View file

@ -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:

View file

@ -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

View file

@ -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.

View file

@ -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<{
</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="trained_heuristic" className="mt-0.5" disabled={scorerLocked} />
<span>
<strong className="font-semibold">Trained heuristic</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>

View file

@ -127,6 +127,31 @@ describe("ComplexityRouterConfig", () => {
expect(onChange).toHaveBeenCalledWith(expectedValue);
});
it("selects the trained heuristic 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("Trained heuristic"));
expect(onChange).toHaveBeenCalledWith(
expect.objectContaining({
classifier_type: "trained_heuristic",
classifier_llm_config: undefined,
}),
);
const trainedValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "trained_heuristic" };
rerender(<ComplexityRouterConfig modelInfo={mockModelInfo} value={trainedValue} 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,

View file

@ -128,7 +128,7 @@ export interface ClassifierLLMConfig {
system_prompt?: string;
}
export type ClassifierType = "heuristic" | "llm" | "heuristic_first";
export type ClassifierType = "heuristic" | "trained_heuristic" | "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 === "trained_heuristic") 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 === "trained_heuristic") {
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 ??

View file

@ -165,6 +165,12 @@ describe("ClassificationMethodConfig scorer gating", () => {
it.each([
["heuristic decides the tier", "heuristic" as ClassifierType, undefined, true],
[
"the trained heuristic decides without the weighted scorer",
"trained_heuristic" 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",

View file

@ -111,6 +111,21 @@ describe("buildComplexityRouterConfig", () => {
expect(config.classifier_llm_config).toBeUndefined();
});
it("emits trained_heuristic without classifier-only fields", () => {
const trainedParams: BuildComplexityRouterConfigParams = {
...baseParams,
classifierType: "trained_heuristic",
classifierLlmConfig: { model: "gpt-4o-mini", timeout_ms: 3000 },
classifierContextWindowSize: 5,
classifierFallback: "heuristic",
};
const config = buildComplexityRouterConfig(trainedParams);
expect(config.classifier_type).toBe("trained_heuristic");
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 trained heuristic router, which runs locally", () => {
expect(getClassifierModelError({ classifier_type: "trained_heuristic" })).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", "trained_heuristic", "llm"] as const) {
const config = buildComplexityRouterConfig({ ...heuristicFirstParams, classifierType });
expect(config.heuristic_first_max_tier).toBeUndefined();
}

View file

@ -84,6 +84,7 @@ function describeReasoningOverride(tierLabel: string | undefined, floor: number
const CONSTANT_CAUSE_LABELS: Record<string, string> = {
heuristic_scorer: "Heuristic scorer",
trained_heuristic: "Trained tier heuristic",
heuristic_first_short_circuit: "Heuristic scorer, classifier skipped",
classifier_plugin: "Custom classifier plugin",
semantic_keyword_match: "Semantic keyword match",

View file

@ -22696,6 +22696,103 @@ export interface components {
/** Results */
results: components["schemas"]["TagActiveUsersResponse"][];
};
/** AdaptiveRouterTierArtifact */
AdaptiveRouterTierArtifact: {
/**
* Cohort Prior Mass
* @default 20
*/
cohort_prior_mass: number;
/**
* Cohort Statistics
* @default []
*/
cohort_statistics: components["schemas"]["AdaptiveRouterTierCohortStatistic"][];
/**
* Datasets
* @default []
*/
datasets: components["schemas"]["AdaptiveRouterTierDataset"][];
/**
* Domain Prior Mass
* @default 200
*/
domain_prior_mass: number;
/**
* Domain Statistics
* @default []
*/
domain_statistics: components["schemas"]["AdaptiveRouterTierDomainStatistic"][];
/** Global Statistics */
global_statistics: components["schemas"]["AdaptiveRouterTierGlobalStatistic"][];
/**
* 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;
};
/** AdaptiveRouterTierCohortStatistic */
AdaptiveRouterTierCohortStatistic: {
/** Cohort */
cohort: string;
/** Observations */
observations: number;
/** Successes */
successes: number;
/** Tier */
tier: number;
};
/** AdaptiveRouterTierDataset */
AdaptiveRouterTierDataset: {
/** 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;
};
/** AdaptiveRouterTierDomainStatistic */
AdaptiveRouterTierDomainStatistic: {
/** Observations */
observations: number;
request_type: components["schemas"]["RequestType"];
/** Successes */
successes: number;
/** Tier */
tier: number;
};
/** AdaptiveRouterTierGlobalStatistic */
AdaptiveRouterTierGlobalStatistic: {
/** Observations */
observations: number;
/** Successes */
successes: number;
/** Tier */
tier: number;
};
/** AdaptiveRouterWeights */
AdaptiveRouterWeights: {
/**
@ -34444,11 +34541,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" | "trained_heuristic" | "llm" | "custom" | "heuristic_first";
/**
* Code Keywords
* @description Keywords indicating code-related content
@ -34644,9 +34741,21 @@ export interface components {
token_thresholds?: {
[key: string]: number;
};
/**
* Trained Heuristic Artifact
* @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
* @default ultrafeedback
*/
trained_heuristic_artifact: components["schemas"]["AdaptiveRouterTierArtifact"] | "ultrafeedback";
} & {
[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 */
@ -35724,7 +35833,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" | "trained_heuristic" | "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 */