mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(router_strategy): add hybrid router combining complexity gating with per-tier MAB selection
Two-layer routing strategy: Layer 1 classifies request complexity into SIMPLE/MEDIUM/COMPLEX/REASONING tiers using heuristic scoring, Layer 2 picks the best model within each tier using a per-tier multi-armed bandit (Thompson Sampling, UCB, or Epsilon-Greedy). Per-tier bandits learn independently, so a model that is fast for simple queries but slow for complex ones gets scored separately at each tier. Includes eval script and 50 unit tests covering all three bandit algorithms, classification, reward recording, convergence, and tier priors.
This commit is contained in:
parent
8159f240c4
commit
3ae085c639
6 changed files with 2849 additions and 0 deletions
6
litellm/router_strategy/hybrid_router/__init__.py
Normal file
6
litellm/router_strategy/hybrid_router/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from litellm.router_strategy.hybrid_router.hybrid_router import (
|
||||
HybridRouter,
|
||||
HybridRouterConfig,
|
||||
)
|
||||
|
||||
__all__ = ["HybridRouter", "HybridRouterConfig"]
|
||||
252
litellm/router_strategy/hybrid_router/bandit.py
Normal file
252
litellm/router_strategy/hybrid_router/bandit.py
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
"""
|
||||
Multi-armed bandit algorithms for the hybrid router.
|
||||
|
||||
Three algorithms ship: UCB, EpsilonGreedy, and ThompsonSampling.
|
||||
ThompsonSampling reuses the BanditCell and thompson_sample from the
|
||||
adaptive_router module. All three expose the same interface so the
|
||||
hybrid router can swap algorithms via config.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import random
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.router_strategy.adaptive_router.bandit import (
|
||||
BanditCell,
|
||||
apply_delta,
|
||||
thompson_sample,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ArmStats:
|
||||
"""Frequentist statistics for one arm (UCB / EpsilonGreedy)."""
|
||||
|
||||
count: int
|
||||
sum_rewards: float
|
||||
|
||||
@property
|
||||
def mean(self) -> float:
|
||||
return self.sum_rewards / self.count if self.count > 0 else 0.0
|
||||
|
||||
|
||||
class MABAlgorithm(ABC):
|
||||
"""Thread-safe bandit base class."""
|
||||
|
||||
def __init__(self, arms: tuple[str, ...]) -> None:
|
||||
if len(arms) < 2:
|
||||
raise ValueError("Need at least 2 arms")
|
||||
self._arms: Final = arms
|
||||
self._lock: Final = threading.Lock()
|
||||
|
||||
@property
|
||||
def arms(self) -> tuple[str, ...]:
|
||||
return self._arms
|
||||
|
||||
@abstractmethod
|
||||
def select_arm(self) -> str: ...
|
||||
|
||||
@abstractmethod
|
||||
def select_arm_from(self, eligible: tuple[str, ...]) -> str: ...
|
||||
|
||||
@abstractmethod
|
||||
def update(self, arm: str, reward: float) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def state(self) -> dict[str, object]: ...
|
||||
|
||||
|
||||
class UCB(MABAlgorithm):
|
||||
"""Upper Confidence Bound (Hoeffding-style).
|
||||
|
||||
Select argmax(mu_i + sqrt(2 * ln(1/delta) / n_i)).
|
||||
Smaller delta = wider interval = more exploration.
|
||||
"""
|
||||
|
||||
def __init__(self, arms: tuple[str, ...], *, delta: float = 0.1) -> None:
|
||||
super().__init__(arms)
|
||||
if not (0.0 < delta < 1.0):
|
||||
raise ValueError("delta must be in (0, 1)")
|
||||
self._delta: Final = delta
|
||||
self._stats: dict[str, ArmStats] = {arm: ArmStats(count=0, sum_rewards=0.0) for arm in arms}
|
||||
self._t = 0
|
||||
|
||||
def _ucb_score(self, arm: str) -> float | None:
|
||||
stats: Final = self._stats[arm]
|
||||
if stats.count == 0:
|
||||
return None
|
||||
return stats.mean + math.sqrt(2.0 * math.log(1.0 / self._delta) / stats.count)
|
||||
|
||||
def select_arm(self) -> str:
|
||||
with self._lock:
|
||||
return self._pick(self._arms)
|
||||
|
||||
def select_arm_from(self, eligible: tuple[str, ...]) -> str:
|
||||
with self._lock:
|
||||
return self._pick(eligible)
|
||||
|
||||
def _pick(self, candidates: tuple[str, ...]) -> str:
|
||||
best_arm: str = candidates[0]
|
||||
best_score: float = -math.inf
|
||||
for arm in candidates:
|
||||
stats: Final = self._stats[arm]
|
||||
if stats.count == 0:
|
||||
return arm
|
||||
score: Final = stats.mean + math.sqrt(2.0 * math.log(1.0 / self._delta) / stats.count)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_arm = arm
|
||||
return best_arm
|
||||
|
||||
def update(self, arm: str, reward: float) -> None:
|
||||
with self._lock:
|
||||
old: Final = self._stats[arm]
|
||||
self._stats[arm] = ArmStats(count=old.count + 1, sum_rewards=old.sum_rewards + reward)
|
||||
self._t += 1
|
||||
|
||||
def state(self) -> dict[str, object]:
|
||||
with self._lock:
|
||||
return {
|
||||
"algorithm": "UCB",
|
||||
"arms": list(self._arms),
|
||||
"t": self._t,
|
||||
"counts": {arm: self._stats[arm].count for arm in self._arms},
|
||||
"means": {arm: self._stats[arm].mean for arm in self._arms},
|
||||
}
|
||||
|
||||
|
||||
class EpsilonGreedy(MABAlgorithm):
|
||||
"""Epsilon-greedy: explore uniformly with probability epsilon, else greedy."""
|
||||
|
||||
def __init__(self, arms: tuple[str, ...], *, epsilon: float = 0.1) -> None:
|
||||
super().__init__(arms)
|
||||
if not (0.0 <= epsilon <= 1.0):
|
||||
raise ValueError("epsilon must be in [0, 1]")
|
||||
self._epsilon: Final = epsilon
|
||||
self._stats: dict[str, ArmStats] = {arm: ArmStats(count=0, sum_rewards=0.0) for arm in arms}
|
||||
self._t = 0
|
||||
|
||||
def select_arm(self) -> str:
|
||||
with self._lock:
|
||||
return self._pick(self._arms)
|
||||
|
||||
def select_arm_from(self, eligible: tuple[str, ...]) -> str:
|
||||
with self._lock:
|
||||
return self._pick(eligible)
|
||||
|
||||
def _pick(self, candidates: tuple[str, ...]) -> str:
|
||||
for arm in candidates:
|
||||
if self._stats[arm].count == 0:
|
||||
return arm
|
||||
if random.random() < self._epsilon:
|
||||
return random.choice(candidates)
|
||||
best_arm: str = candidates[0]
|
||||
best_mean: float = -math.inf
|
||||
for arm in candidates:
|
||||
mean: Final = self._stats[arm].mean
|
||||
if mean > best_mean:
|
||||
best_mean = mean
|
||||
best_arm = arm
|
||||
return best_arm
|
||||
|
||||
def update(self, arm: str, reward: float) -> None:
|
||||
with self._lock:
|
||||
old: Final = self._stats[arm]
|
||||
self._stats[arm] = ArmStats(count=old.count + 1, sum_rewards=old.sum_rewards + reward)
|
||||
self._t += 1
|
||||
|
||||
def state(self) -> dict[str, object]:
|
||||
with self._lock:
|
||||
return {
|
||||
"algorithm": "EpsilonGreedy",
|
||||
"arms": list(self._arms),
|
||||
"t": self._t,
|
||||
"counts": {arm: self._stats[arm].count for arm in self._arms},
|
||||
"means": {arm: self._stats[arm].mean for arm in self._arms},
|
||||
}
|
||||
|
||||
|
||||
class ThompsonSampling(MABAlgorithm):
|
||||
"""Beta-Bernoulli Thompson Sampling.
|
||||
|
||||
Reuses BanditCell and thompson_sample from the adaptive_router module.
|
||||
Rewards in [0, 1] are treated as fractional: alpha += reward, beta += (1 - reward).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
arms: tuple[str, ...],
|
||||
*,
|
||||
arm_priors: Mapping[str, tuple[float, float]] | None = None,
|
||||
prior_alpha: float = 1.0,
|
||||
prior_beta: float = 1.0,
|
||||
) -> None:
|
||||
super().__init__(arms)
|
||||
self._cells: dict[str, BanditCell] = {}
|
||||
for arm in arms:
|
||||
if arm_priors and arm in arm_priors:
|
||||
alpha, beta = arm_priors[arm]
|
||||
else:
|
||||
alpha, beta = prior_alpha, prior_beta
|
||||
self._cells[arm] = BanditCell(alpha=alpha, beta=beta)
|
||||
self._t = 0
|
||||
|
||||
def select_arm(self) -> str:
|
||||
with self._lock:
|
||||
return self._pick(self._arms)
|
||||
|
||||
def select_arm_from(self, eligible: tuple[str, ...]) -> str:
|
||||
with self._lock:
|
||||
return self._pick(eligible)
|
||||
|
||||
def _pick(self, candidates: tuple[str, ...]) -> str:
|
||||
best_arm: str = candidates[0]
|
||||
best_sample: float = -1.0
|
||||
for arm in candidates:
|
||||
sample: Final = thompson_sample(self._cells[arm])
|
||||
if sample > best_sample:
|
||||
best_sample = sample
|
||||
best_arm = arm
|
||||
return best_arm
|
||||
|
||||
def update(self, arm: str, reward: float) -> None:
|
||||
with self._lock:
|
||||
cell: Final = self._cells[arm]
|
||||
self._cells[arm] = BanditCell(alpha=cell.alpha + reward, beta=cell.beta + (1.0 - reward))
|
||||
self._t += 1
|
||||
|
||||
def state(self) -> dict[str, object]:
|
||||
with self._lock:
|
||||
return {
|
||||
"algorithm": "ThompsonSampling",
|
||||
"arms": list(self._arms),
|
||||
"t": self._t,
|
||||
"counts": {arm: int(self._cells[arm].alpha + self._cells[arm].beta - 2.0) for arm in self._arms},
|
||||
"means": {arm: self._cells[arm].mean for arm in self._arms},
|
||||
"alpha": {arm: self._cells[arm].alpha for arm in self._arms},
|
||||
"beta": {arm: self._cells[arm].beta for arm in self._arms},
|
||||
}
|
||||
|
||||
|
||||
BANDIT_REGISTRY: Final[Mapping[str, Callable[..., MABAlgorithm]]] = MappingProxyType(
|
||||
{
|
||||
"ucb": lambda arms, *, delta=0.1, **_kw: UCB(arms, delta=delta),
|
||||
"epsilon_greedy": lambda arms, *, epsilon=0.1, **_kw: EpsilonGreedy(arms, epsilon=epsilon),
|
||||
"thompson": lambda arms, *, arm_priors=None, prior_alpha=1.0, prior_beta=1.0, **_kw: ThompsonSampling(
|
||||
arms, arm_priors=arm_priors, prior_alpha=prior_alpha, prior_beta=prior_beta
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def make_bandit(name: str, arms: tuple[str, ...], **kwargs: object) -> MABAlgorithm:
|
||||
if name not in BANDIT_REGISTRY:
|
||||
raise ValueError(f"Unknown bandit {name!r}. Available: {sorted(BANDIT_REGISTRY)}")
|
||||
return BANDIT_REGISTRY[name](arms, **kwargs)
|
||||
318
litellm/router_strategy/hybrid_router/hybrid_router.py
Normal file
318
litellm/router_strategy/hybrid_router/hybrid_router.py
Normal file
|
|
@ -0,0 +1,318 @@
|
|||
"""
|
||||
Hybrid Router: two-layer routing combining complexity gating with online efficiency.
|
||||
|
||||
Layer 1 (Complexity): classifies the request using the same heuristic scorer as the
|
||||
complexity router and produces a set of eligible models from the tier's candidate pool.
|
||||
|
||||
Layer 2 (Efficiency): picks the best model from that set using a per-tier MAB bandit
|
||||
that optimizes for latency, throughput, and optionally quality in real time.
|
||||
|
||||
The per-tier bandit design means each tier tracks arm performance independently.
|
||||
A model might be fast for SIMPLE queries but slow for COMPLEX ones (longer generation),
|
||||
and the bandits learn this separately.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import threading
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.router_strategy.complexity_router.config import (
|
||||
DEFAULT_CODE_KEYWORDS,
|
||||
DEFAULT_DIMENSION_WEIGHTS,
|
||||
DEFAULT_REASONING_KEYWORDS,
|
||||
DEFAULT_SIMPLE_KEYWORDS,
|
||||
DEFAULT_TECHNICAL_KEYWORDS,
|
||||
DEFAULT_TIER_BOUNDARIES,
|
||||
DEFAULT_TOKEN_THRESHOLDS,
|
||||
ComplexityTier,
|
||||
)
|
||||
from litellm.router_strategy.hybrid_router.bandit import MABAlgorithm, make_bandit
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class HybridRouterConfig:
|
||||
tier_candidates: Mapping[str, tuple[str, ...]]
|
||||
bandit: str = "thompson"
|
||||
delta: float = 0.1
|
||||
epsilon: float = 0.1
|
||||
target_tpt: float = 1.0
|
||||
tier_priors: Mapping[str, Mapping[str, tuple[float, float]]] | None = None
|
||||
dimension_weights: Mapping[str, float] | None = None
|
||||
tier_boundaries: Mapping[str, float] | None = None
|
||||
token_thresholds: Mapping[str, int] | None = None
|
||||
code_keywords: tuple[str, ...] | None = None
|
||||
reasoning_keywords: tuple[str, ...] | None = None
|
||||
technical_keywords: tuple[str, ...] | None = None
|
||||
simple_keywords: tuple[str, ...] | None = None
|
||||
|
||||
|
||||
class TierBandit:
|
||||
"""A bandit instance scoped to one complexity tier."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tier: str,
|
||||
arms: tuple[str, ...],
|
||||
bandit_type: str,
|
||||
arm_priors: Mapping[str, tuple[float, float]] | None = None,
|
||||
**bandit_kwargs: object,
|
||||
) -> None:
|
||||
self._tier: Final = tier
|
||||
self._arms: Final = arms
|
||||
self._bandit: Final[MABAlgorithm] = make_bandit(
|
||||
bandit_type, arms, arm_priors=arm_priors, **bandit_kwargs
|
||||
)
|
||||
|
||||
@property
|
||||
def tier(self) -> str:
|
||||
return self._tier
|
||||
|
||||
def pick(self) -> str:
|
||||
return self._bandit.select_arm()
|
||||
|
||||
def pick_from(self, eligible: tuple[str, ...]) -> str:
|
||||
if len(eligible) == 1:
|
||||
return eligible[0]
|
||||
return self._bandit.select_arm_from(eligible)
|
||||
|
||||
def update(self, arm: str, reward: float) -> None:
|
||||
self._bandit.update(arm, reward)
|
||||
|
||||
def state(self) -> dict[str, object]:
|
||||
return {"tier": self._tier, **self._bandit.state()}
|
||||
|
||||
|
||||
_MULTI_STEP_PATTERNS: Final = (
|
||||
re.compile(r"first.*?then", re.IGNORECASE),
|
||||
re.compile(r"step\s*\d", re.IGNORECASE),
|
||||
re.compile(r"\d+\.\s"),
|
||||
re.compile(r"[a-z]\)\s", re.IGNORECASE),
|
||||
)
|
||||
|
||||
|
||||
class HybridRouter:
|
||||
"""
|
||||
Two-layer router: complexity classification -> per-tier MAB selection.
|
||||
|
||||
Usage:
|
||||
config = HybridRouterConfig(
|
||||
tier_candidates={
|
||||
"SIMPLE": ("model-small", "model-medium"),
|
||||
"MEDIUM": ("model-medium", "model-large"),
|
||||
"COMPLEX": ("model-large",),
|
||||
"REASONING": ("model-large",),
|
||||
},
|
||||
bandit="thompson",
|
||||
)
|
||||
router = HybridRouter(config)
|
||||
|
||||
# On each request:
|
||||
tier, model = router.route(user_message, system_prompt)
|
||||
|
||||
# After response:
|
||||
router.record_success(tier, model, latency, completion_tokens)
|
||||
"""
|
||||
|
||||
def __init__(self, config: HybridRouterConfig) -> None:
|
||||
self._config: Final = config
|
||||
self._tier_candidates: Final = config.tier_candidates
|
||||
|
||||
all_models: Final = tuple(
|
||||
sorted(frozenset(model for models in self._tier_candidates.values() for model in models))
|
||||
)
|
||||
self._all_models: Final = all_models
|
||||
|
||||
bandit_kwargs: Final[dict[str, object]] = {"delta": config.delta, "epsilon": config.epsilon}
|
||||
tier_bandits: Final[dict[str, TierBandit]] = {}
|
||||
for tier_name, candidates in self._tier_candidates.items():
|
||||
if len(candidates) >= 2:
|
||||
arm_priors = (config.tier_priors or {}).get(tier_name)
|
||||
tier_bandits[tier_name] = TierBandit(
|
||||
tier=tier_name,
|
||||
arms=candidates,
|
||||
bandit_type=config.bandit,
|
||||
arm_priors=arm_priors,
|
||||
**bandit_kwargs,
|
||||
)
|
||||
self._tier_bandits: Final = tier_bandits
|
||||
|
||||
self._code_keywords: Final = config.code_keywords or tuple(DEFAULT_CODE_KEYWORDS)
|
||||
self._reasoning_keywords: Final = config.reasoning_keywords or tuple(DEFAULT_REASONING_KEYWORDS)
|
||||
self._technical_keywords: Final = config.technical_keywords or tuple(DEFAULT_TECHNICAL_KEYWORDS)
|
||||
self._simple_keywords: Final = config.simple_keywords or tuple(DEFAULT_SIMPLE_KEYWORDS)
|
||||
self._dimension_weights: Final = config.dimension_weights or MappingProxyType(DEFAULT_DIMENSION_WEIGHTS)
|
||||
self._tier_boundaries: Final = config.tier_boundaries or MappingProxyType(DEFAULT_TIER_BOUNDARIES)
|
||||
self._token_thresholds: Final = config.token_thresholds or MappingProxyType(DEFAULT_TOKEN_THRESHOLDS)
|
||||
|
||||
self._lock: Final = threading.Lock()
|
||||
|
||||
@property
|
||||
def config(self) -> HybridRouterConfig:
|
||||
return self._config
|
||||
|
||||
def classify(
|
||||
self, user_message: str, system_prompt: str | None = None
|
||||
) -> tuple[ComplexityTier, float, tuple[str, ...]]:
|
||||
user_text: Final = user_message.lower()
|
||||
estimated_tokens: Final = len(user_message) // 4
|
||||
|
||||
simple_threshold: Final = self._token_thresholds.get("simple", 15)
|
||||
complex_threshold: Final = self._token_thresholds.get("complex", 400)
|
||||
|
||||
scores: dict[str, float] = {}
|
||||
signals: list[str] = []
|
||||
|
||||
if estimated_tokens < simple_threshold:
|
||||
scores["tokenCount"] = -1.0
|
||||
elif estimated_tokens > complex_threshold:
|
||||
scores["tokenCount"] = 1.0
|
||||
else:
|
||||
scores["tokenCount"] = 0.0
|
||||
|
||||
def keyword_score(
|
||||
text: str,
|
||||
keywords: tuple[str, ...],
|
||||
name: str,
|
||||
label: str,
|
||||
thresholds: tuple[int, int],
|
||||
score_vals: tuple[float, float, float],
|
||||
) -> int:
|
||||
matches: Final = tuple(
|
||||
kw for kw in keywords if _keyword_matches(text, kw)
|
||||
)
|
||||
count: Final = len(matches)
|
||||
low_t, high_t = thresholds
|
||||
if count >= high_t:
|
||||
scores[name] = score_vals[2]
|
||||
signals.append(f"{label} ({', '.join(matches[:3])})")
|
||||
elif count >= low_t:
|
||||
scores[name] = score_vals[1]
|
||||
signals.append(f"{label} ({', '.join(matches[:3])})")
|
||||
else:
|
||||
scores[name] = score_vals[0]
|
||||
return count
|
||||
|
||||
keyword_score(user_text, self._code_keywords, "codePresence", "code", (1, 2), (0, 0.5, 1.0))
|
||||
reasoning_count: Final = keyword_score(
|
||||
user_text, self._reasoning_keywords, "reasoningMarkers", "reasoning", (1, 2), (0, 0.7, 1.0)
|
||||
)
|
||||
keyword_score(user_text, self._technical_keywords, "technicalTerms", "technical", (2, 4), (0, 0.5, 1.0))
|
||||
keyword_score(user_text, self._simple_keywords, "simpleIndicators", "simple", (1, 2), (0, -1.0, -1.0))
|
||||
|
||||
multi_step_hits: Final = sum(1 for p in _MULTI_STEP_PATTERNS if p.search(user_text))
|
||||
scores["multiStepPatterns"] = 0.5 if multi_step_hits > 0 else 0.0
|
||||
|
||||
q_count: Final = user_message.count("?")
|
||||
scores["questionComplexity"] = 0.5 if q_count > 3 else 0.0
|
||||
|
||||
weighted_score: Final = sum(scores.get(name, 0) * w for name, w in self._dimension_weights.items())
|
||||
|
||||
if reasoning_count >= 2:
|
||||
return ComplexityTier.REASONING, weighted_score, tuple(signals)
|
||||
|
||||
simple_medium: Final = self._tier_boundaries.get("simple_medium", 0.15)
|
||||
medium_complex: Final = self._tier_boundaries.get("medium_complex", 0.35)
|
||||
complex_reasoning: Final = self._tier_boundaries.get("complex_reasoning", 0.60)
|
||||
|
||||
if weighted_score < simple_medium:
|
||||
tier = ComplexityTier.SIMPLE
|
||||
elif weighted_score < medium_complex:
|
||||
tier = ComplexityTier.MEDIUM
|
||||
elif weighted_score < complex_reasoning:
|
||||
tier = ComplexityTier.COMPLEX
|
||||
else:
|
||||
tier = ComplexityTier.REASONING
|
||||
|
||||
return tier, weighted_score, tuple(signals)
|
||||
|
||||
def get_candidates(self, tier: ComplexityTier) -> tuple[str, ...]:
|
||||
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
|
||||
candidates = self._tier_candidates.get(tier_key)
|
||||
if candidates:
|
||||
return candidates
|
||||
for fallback_key in ("MEDIUM", "COMPLEX", "SIMPLE", "REASONING"):
|
||||
if fallback_key in self._tier_candidates:
|
||||
return self._tier_candidates[fallback_key]
|
||||
return self._all_models
|
||||
|
||||
def pick_model(self, tier: ComplexityTier) -> str:
|
||||
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
|
||||
candidates: Final = self.get_candidates(tier)
|
||||
|
||||
if len(candidates) == 1:
|
||||
return candidates[0]
|
||||
|
||||
bandit = self._tier_bandits.get(tier_key)
|
||||
if bandit is None:
|
||||
return candidates[0]
|
||||
|
||||
return bandit.pick_from(candidates)
|
||||
|
||||
def route(self, user_message: str, system_prompt: str | None = None) -> tuple[ComplexityTier, str]:
|
||||
tier, _score, _signals = self.classify(user_message, system_prompt)
|
||||
model: Final = self.pick_model(tier)
|
||||
return tier, model
|
||||
|
||||
def _compute_reward(self, latency_seconds: float, completion_tokens: int) -> float:
|
||||
time_per_token: Final = latency_seconds / max(completion_tokens, 1)
|
||||
return self._config.target_tpt / (self._config.target_tpt + time_per_token)
|
||||
|
||||
def record_success(
|
||||
self, tier: ComplexityTier, model: str, latency_seconds: float, completion_tokens: int = 0
|
||||
) -> float:
|
||||
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
|
||||
reward: Final = self._compute_reward(latency_seconds, completion_tokens)
|
||||
bandit = self._tier_bandits.get(tier_key)
|
||||
if bandit is not None:
|
||||
bandit.update(model, reward)
|
||||
return reward
|
||||
|
||||
def record_quality(
|
||||
self,
|
||||
tier: ComplexityTier,
|
||||
model: str,
|
||||
latency_seconds: float,
|
||||
completion_tokens: int,
|
||||
quality_score: float,
|
||||
quality_weight: float = 0.5,
|
||||
) -> float:
|
||||
"""Record a composite reward blending efficiency and quality.
|
||||
|
||||
quality_score: Judge/quality score in [0, 1].
|
||||
quality_weight: How much to weight quality vs efficiency. 0.5 = equal.
|
||||
"""
|
||||
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
|
||||
eff_reward: Final = self._compute_reward(latency_seconds, completion_tokens)
|
||||
composite: Final = (1.0 - quality_weight) * eff_reward + quality_weight * quality_score
|
||||
bandit = self._tier_bandits.get(tier_key)
|
||||
if bandit is not None:
|
||||
bandit.update(model, composite)
|
||||
return composite
|
||||
|
||||
def record_failure(self, tier: ComplexityTier, model: str) -> None:
|
||||
tier_key: Final = tier.value if isinstance(tier, ComplexityTier) else tier
|
||||
bandit = self._tier_bandits.get(tier_key)
|
||||
if bandit is not None:
|
||||
bandit.update(model, 0.0)
|
||||
|
||||
def state(self) -> dict[str, object]:
|
||||
return {
|
||||
"config": {
|
||||
"bandit": self._config.bandit,
|
||||
"tier_candidates": {k: list(v) for k, v in self._config.tier_candidates.items()},
|
||||
"target_tpt": self._config.target_tpt,
|
||||
},
|
||||
"tier_bandits": {tier: bandit.state() for tier, bandit in self._tier_bandits.items()},
|
||||
}
|
||||
|
||||
|
||||
def _keyword_matches(text: str, keyword: str) -> bool:
|
||||
kw_lower: Final = keyword.lower()
|
||||
if " " in kw_lower:
|
||||
return kw_lower in text
|
||||
return bool(re.search(r"\b" + re.escape(kw_lower) + r"\b", text))
|
||||
1179
scripts/efficiency_eval_prompts.py
Normal file
1179
scripts/efficiency_eval_prompts.py
Normal file
File diff suppressed because it is too large
Load diff
717
scripts/test_custom_router_hybrid_v1.py
Normal file
717
scripts/test_custom_router_hybrid_v1.py
Normal file
|
|
@ -0,0 +1,717 @@
|
|||
"""
|
||||
Hybrid Router v1.0 Evaluation.
|
||||
|
||||
Thompson Sampling with session-based judge feedback. Runs prompts in sessions
|
||||
of BATCH_SIZE. After each session, judges all responses, then feeds composite
|
||||
(efficiency + quality) rewards into the Thompson Sampling bandit. This lets
|
||||
the bandit learn per-tier model preferences that balance speed and correctness.
|
||||
|
||||
Assumes:
|
||||
- Qwen3.5-9B served at http://localhost:8001/v1
|
||||
- Qwen3.5-4B served at http://localhost:8002/v1
|
||||
- CUDA_VISIBLE_DEVICES=0 vllm serve Qwen/Qwen3.5-9B --port 8001 --served-model-name Qwen3.5-9b --gpu-memory-utilization 0.90
|
||||
- CUDA_VISIBLE_DEVICES=1 vllm serve Qwen/Qwen3.5-4B --port 8002 --served-model-name Qwen3.5-4b --gpu-memory-utilization 0.90
|
||||
Usage:
|
||||
python scripts/test_custom_router_euro_v1.py
|
||||
|
||||
Environment variables:
|
||||
VLLM_9B_BASE - 9B model endpoint (default: http://localhost:8001/v1)
|
||||
VLLM_4B_BASE - 4B model endpoint (default: http://localhost:8002/v1)
|
||||
HYBRID_EVAL_CONCURRENCY - Concurrent inference requests (default: 16)
|
||||
JUDGE_CONCURRENCY - Concurrent judge requests (default: 32)
|
||||
JUDGE_TIMEOUT - Judge request timeout in seconds (default: 15)
|
||||
JUDGE_SAMPLE_SIZE - Number of prompts to judge (default: 100)
|
||||
BATCH_SIZE - Prompts per session (default: 200)
|
||||
QUALITY_WEIGHT - Quality vs efficiency weight, 0-1 (default: 0.5)
|
||||
SKIP_JUDGE - Set to "1" to skip accuracy evaluation
|
||||
QUICK_MODE - Set to "1" for 5 sessions of 20 prompts (100 total)
|
||||
RUN_TIER_STATIC - Set to "1" to run tier-static baseline comparison
|
||||
RUN_ROUND_ROBIN - Set to "1" to run round-robin baseline after the main eval
|
||||
CPU_INFERENCE - Set to "1" when vLLM serve is hosted on CPU; runs batched inference mode
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.router_strategy.hybrid_router import HybridRouter, HybridRouterConfig
|
||||
|
||||
litellm.set_verbose = False
|
||||
|
||||
VLLM_9B_BASE = os.getenv("VLLM_9B_BASE", "http://localhost:8001/v1")
|
||||
VLLM_4B_BASE = os.getenv("VLLM_4B_BASE", "http://localhost:8002/v1")
|
||||
|
||||
MODEL_9B = "Qwen3.5-9b"
|
||||
MODEL_4B = "Qwen3.5-4b"
|
||||
|
||||
JUDGE_MODEL = "openai/google.gemma-3-12b-it"
|
||||
JUDGE_CONCURRENCY = int(os.getenv("JUDGE_CONCURRENCY", "32"))
|
||||
JUDGE_TIMEOUT = float(os.getenv("JUDGE_TIMEOUT", "15"))
|
||||
JUDGE_SAMPLE_SIZE = int(os.getenv("JUDGE_SAMPLE_SIZE", "100"))
|
||||
|
||||
CONCURRENCY = int(os.getenv("HYBRID_EVAL_CONCURRENCY", "16"))
|
||||
QUALITY_WEIGHT = float(os.getenv("QUALITY_WEIGHT", "0.5"))
|
||||
|
||||
_quick = os.getenv("QUICK_MODE", "0") == "1"
|
||||
BATCH_SIZE = int(os.getenv("BATCH_SIZE", "20" if _quick else "200"))
|
||||
MAX_PROMPTS = 100 if _quick else None
|
||||
RUN_TIER_STATIC = os.getenv("RUN_TIER_STATIC", "0") == "1"
|
||||
RUN_ROUND_ROBIN = os.getenv("RUN_ROUND_ROBIN", "0") == "1"
|
||||
CPU_INFERENCE = os.getenv("CPU_INFERENCE", "0") == "1"
|
||||
|
||||
CALIBRATION_PROMPT = "What is 2 + 2?"
|
||||
CALIBRATION_MAX_TOKENS = 64
|
||||
|
||||
|
||||
async def calibrate_target_tpt(router: Router, smallest_model: str) -> float:
|
||||
t0 = time.perf_counter()
|
||||
resp = await router.acompletion(
|
||||
model=smallest_model,
|
||||
messages=[{"role": "user", "content": CALIBRATION_PROMPT}],
|
||||
max_tokens=CALIBRATION_MAX_TOKENS,
|
||||
)
|
||||
latency = time.perf_counter() - t0
|
||||
tokens = resp.usage.completion_tokens if resp.usage else CALIBRATION_MAX_TOKENS
|
||||
tpt = latency / max(tokens, 1)
|
||||
print(f" Calibration ({smallest_model}): latency={latency:.2f}s, tokens={tokens}, tpt={tpt:.3f}s/tok")
|
||||
return tpt
|
||||
|
||||
|
||||
async def main():
|
||||
from scripts.efficiency_eval_prompts import PROMPT_METADATA
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "qwen-local-9b",
|
||||
"litellm_params": {
|
||||
"model": f"openai/{MODEL_9B}",
|
||||
"api_base": VLLM_9B_BASE,
|
||||
"api_key": "not-needed",
|
||||
},
|
||||
"model_info": {
|
||||
"id": "Qwen-9b",
|
||||
"cache_creation_input_token_cost": 0.0,
|
||||
"cache_read_input_token_cost": 0.0,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "qwen-local-4b",
|
||||
"litellm_params": {
|
||||
"model": f"openai/{MODEL_4B}",
|
||||
"api_base": VLLM_4B_BASE,
|
||||
"api_key": "not-needed",
|
||||
},
|
||||
"model_info": {
|
||||
"id": "Qwen-4b",
|
||||
"cache_creation_input_token_cost": 0.0,
|
||||
"cache_read_input_token_cost": 0.0,
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
print("Calibrating target_tpt...")
|
||||
target_tpt = await calibrate_target_tpt(router, "qwen-local-4b")
|
||||
print()
|
||||
|
||||
hybrid_config = HybridRouterConfig(
|
||||
tier_candidates={
|
||||
"SIMPLE": ("qwen-local-4b", "qwen-local-9b"),
|
||||
"MEDIUM": ("qwen-local-4b", "qwen-local-9b"),
|
||||
"COMPLEX": ("qwen-local-9b",),
|
||||
"REASONING": ("qwen-local-9b",),
|
||||
},
|
||||
bandit="thompson",
|
||||
target_tpt=target_tpt,
|
||||
tier_priors={
|
||||
"SIMPLE": {"qwen-local-4b": (5.0, 1.0), "qwen-local-9b": (1.0, 2.0)},
|
||||
"MEDIUM": {"qwen-local-4b": (1.0, 1.0), "qwen-local-9b": (2.0, 1.0)},
|
||||
},
|
||||
)
|
||||
hybrid = HybridRouter(hybrid_config)
|
||||
|
||||
prompts = list(PROMPT_METADATA)
|
||||
if MAX_PROMPTS is not None:
|
||||
prompts = prompts[:MAX_PROMPTS]
|
||||
num_sessions = (len(prompts) + BATCH_SIZE - 1) // BATCH_SIZE
|
||||
|
||||
skip_judge = os.getenv("SKIP_JUDGE", "0") == "1"
|
||||
req_semaphore = asyncio.Semaphore(CONCURRENCY)
|
||||
judge_semaphore = asyncio.Semaphore(JUDGE_CONCURRENCY)
|
||||
|
||||
bar_width = 40
|
||||
max_tokens_by_tier = {1: 64, 2: 128, 3: 256, 4: 512, 5: 512}
|
||||
|
||||
print(f"=== Hybrid Router v1.0 ===")
|
||||
print(f" Prompts: {len(prompts)}, Batch size: {BATCH_SIZE}, Sessions: {num_sessions}")
|
||||
print(f" Quality weight: {QUALITY_WEIGHT}")
|
||||
print(f" Judge: {JUDGE_MODEL}, Skip judge: {skip_judge}")
|
||||
print(f" Concurrency: inference={CONCURRENCY}, judge={JUDGE_CONCURRENCY}")
|
||||
mode_label = "tier-static baseline" if RUN_TIER_STATIC else "bandit"
|
||||
if RUN_ROUND_ROBIN:
|
||||
mode_label += " + round-robin baseline"
|
||||
print(f" Mode: {mode_label}")
|
||||
print()
|
||||
|
||||
all_results: list[dict] = []
|
||||
session_summaries: list[dict] = []
|
||||
|
||||
total_steps = len(prompts) * 2
|
||||
completed_steps = 0
|
||||
|
||||
def update_progress(session_idx: int, phase: str) -> None:
|
||||
nonlocal completed_steps
|
||||
completed_steps += 1
|
||||
pct = completed_steps / total_steps
|
||||
filled = int(bar_width * pct)
|
||||
bar = "#" * filled + "-" * (bar_width - filled)
|
||||
print(f"\r [{bar}] {pct*100:.1f}% | Session {session_idx+1}/{num_sessions} ({phase})", end="", flush=True)
|
||||
|
||||
t_total_start = time.perf_counter()
|
||||
|
||||
if RUN_TIER_STATIC:
|
||||
print("BASELINE: Tier-Static (SIMPLE->4B, MEDIUM/COMPLEX/REASONING->9B)\n")
|
||||
|
||||
for session_idx in range(num_sessions):
|
||||
batch_start = session_idx * BATCH_SIZE
|
||||
batch_end = min(batch_start + BATCH_SIZE, len(prompts))
|
||||
batch = prompts[batch_start:batch_end]
|
||||
|
||||
decisions = []
|
||||
for meta in batch:
|
||||
tier, _ = hybrid.route(meta["prompt"])
|
||||
tier_key = tier.value if hasattr(tier, "value") else tier
|
||||
model = "qwen-local-4b" if tier_key == "SIMPLE" else "qwen-local-9b"
|
||||
decisions.append({"tier": tier, "model": model})
|
||||
|
||||
async def run_one(meta, model, _session_idx=session_idx):
|
||||
async with req_semaphore:
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
resp = await router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": meta["prompt"]}],
|
||||
max_tokens=max_tokens_by_tier[meta["tier"]],
|
||||
)
|
||||
latency = time.perf_counter() - t0
|
||||
tokens = resp.usage.completion_tokens if resp.usage else 0
|
||||
text = resp.choices[0].message.content if resp.choices else ""
|
||||
update_progress(_session_idx, "inference")
|
||||
return {"latency": latency, "tokens": tokens, "text": text or "", "success": True}
|
||||
except Exception:
|
||||
update_progress(_session_idx, "inference")
|
||||
return {"latency": time.perf_counter() - t0, "tokens": 0, "text": "", "success": False}
|
||||
|
||||
t0 = time.perf_counter()
|
||||
if CPU_INFERENCE:
|
||||
indexed_batch = list(enumerate(zip(batch, decisions)))
|
||||
responses = [None] * len(indexed_batch)
|
||||
models_in_batch = sorted(set(dec["model"] for dec in decisions))
|
||||
for model in models_in_batch:
|
||||
model_items = [(i, meta, dec) for i, (meta, dec) in indexed_batch if dec["model"] == model]
|
||||
model_responses = await asyncio.gather(*[
|
||||
run_one(meta, dec["model"]) for _, meta, dec in model_items
|
||||
])
|
||||
for (i, _, _), resp in zip(model_items, model_responses):
|
||||
responses[i] = resp
|
||||
else:
|
||||
responses = await asyncio.gather(*[
|
||||
run_one(meta, dec["model"]) for meta, dec in zip(batch, decisions)
|
||||
])
|
||||
session_wall = time.perf_counter() - t0
|
||||
|
||||
if skip_judge:
|
||||
scores = [1] * len(batch)
|
||||
completed_steps += len(batch)
|
||||
update_progress(session_idx, "judge-skip")
|
||||
else:
|
||||
async def judge_one(prompt, text, _session_idx=session_idx):
|
||||
if not text:
|
||||
update_progress(_session_idx, "judging")
|
||||
return 0
|
||||
async with judge_semaphore:
|
||||
judge_prompt = (
|
||||
"You are an accuracy judge. Given a question and an answer, "
|
||||
"determine if the answer is correct.\n\n"
|
||||
"Respond with ONLY a single digit: 1 if the answer is correct, "
|
||||
"0 if it is incorrect or incomplete.\n\n"
|
||||
f"Question: {prompt}\n\nAnswer: {text}\n\nVerdict (1 or 0):"
|
||||
)
|
||||
try:
|
||||
resp = await asyncio.wait_for(
|
||||
litellm.acompletion(
|
||||
model=JUDGE_MODEL,
|
||||
messages=[{"role": "user", "content": judge_prompt}],
|
||||
max_tokens=32,
|
||||
temperature=0.0,
|
||||
),
|
||||
timeout=JUDGE_TIMEOUT,
|
||||
)
|
||||
verdict = resp.choices[0].message.content.strip()
|
||||
update_progress(_session_idx, "judging")
|
||||
return 1 if verdict.startswith("1") else 0
|
||||
except Exception:
|
||||
update_progress(_session_idx, "judging")
|
||||
return -1
|
||||
|
||||
scores = await asyncio.gather(*[
|
||||
judge_one(m["prompt"], r["text"]) for m, r in zip(batch, responses)
|
||||
])
|
||||
|
||||
model_counts: dict[str, int] = {}
|
||||
for meta, dec, resp, score in zip(batch, decisions, responses, scores):
|
||||
model = dec["model"]
|
||||
model_counts[model] = model_counts.get(model, 0) + 1
|
||||
all_results.append({
|
||||
"session": session_idx,
|
||||
"model": model,
|
||||
"tier": dec["tier"].value if hasattr(dec["tier"], "value") else dec["tier"],
|
||||
"category": meta["category"],
|
||||
"difficulty": meta["tier"],
|
||||
"latency": resp["latency"],
|
||||
"tokens": resp["tokens"],
|
||||
"success": resp["success"],
|
||||
"judge_score": score,
|
||||
"reward": 0.0,
|
||||
})
|
||||
|
||||
successful = [r for r in responses if r["success"]]
|
||||
avg_lat = sum(r["latency"] for r in successful) / len(successful) if successful else 0
|
||||
judged = [s for s in scores if s >= 0]
|
||||
accuracy = sum(1 for s in judged if s == 1) / len(judged) if judged else 0
|
||||
|
||||
print(f"\n Session {session_idx+1}/{num_sessions}: "
|
||||
f"models={model_counts}, "
|
||||
f"acc={accuracy:.3f}, "
|
||||
f"lat={avg_lat:.3f}s, "
|
||||
f"wall={session_wall:.1f}s")
|
||||
|
||||
else:
|
||||
for session_idx in range(num_sessions):
|
||||
batch_start = session_idx * BATCH_SIZE
|
||||
batch_end = min(batch_start + BATCH_SIZE, len(prompts))
|
||||
batch = prompts[batch_start:batch_end]
|
||||
|
||||
decisions = []
|
||||
for meta in batch:
|
||||
tier, model = hybrid.route(meta["prompt"])
|
||||
if QUALITY_WEIGHT >= 1.0:
|
||||
model = "qwen-local-9b"
|
||||
elif QUALITY_WEIGHT <= 0.0:
|
||||
model = "qwen-local-4b"
|
||||
decisions.append({"tier": tier, "model": model})
|
||||
|
||||
async def run_one(meta, model, _session_idx=session_idx):
|
||||
async with req_semaphore:
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
resp = await router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": meta["prompt"]}],
|
||||
max_tokens=max_tokens_by_tier[meta["tier"]],
|
||||
)
|
||||
latency = time.perf_counter() - t0
|
||||
tokens = resp.usage.completion_tokens if resp.usage else 0
|
||||
text = resp.choices[0].message.content if resp.choices else ""
|
||||
update_progress(_session_idx, "inference")
|
||||
return {"latency": latency, "tokens": tokens, "text": text or "", "success": True}
|
||||
except Exception:
|
||||
update_progress(_session_idx, "inference")
|
||||
return {"latency": time.perf_counter() - t0, "tokens": 0, "text": "", "success": False}
|
||||
|
||||
t0 = time.perf_counter()
|
||||
if CPU_INFERENCE:
|
||||
indexed_batch = list(enumerate(zip(batch, decisions)))
|
||||
responses = [None] * len(indexed_batch)
|
||||
models_in_batch = sorted(set(dec["model"] for dec in decisions))
|
||||
for model in models_in_batch:
|
||||
model_items = [(i, meta, dec) for i, (meta, dec) in indexed_batch if dec["model"] == model]
|
||||
model_responses = await asyncio.gather(*[
|
||||
run_one(meta, dec["model"]) for _, meta, dec in model_items
|
||||
])
|
||||
for (i, _, _), resp in zip(model_items, model_responses):
|
||||
responses[i] = resp
|
||||
else:
|
||||
responses = await asyncio.gather(*[
|
||||
run_one(meta, dec["model"]) for meta, dec in zip(batch, decisions)
|
||||
])
|
||||
session_wall = time.perf_counter() - t0
|
||||
|
||||
if skip_judge:
|
||||
scores = [1] * len(batch)
|
||||
completed_steps += len(batch)
|
||||
update_progress(session_idx, "judge-skip")
|
||||
else:
|
||||
async def judge_one(prompt, text, _session_idx=session_idx):
|
||||
if not text:
|
||||
update_progress(_session_idx, "judging")
|
||||
return 0
|
||||
async with judge_semaphore:
|
||||
judge_prompt = (
|
||||
"You are an accuracy judge. Given a question and an answer, "
|
||||
"determine if the answer is correct.\n\n"
|
||||
"Respond with ONLY a single digit: 1 if the answer is correct, "
|
||||
"0 if it is incorrect or incomplete.\n\n"
|
||||
f"Question: {prompt}\n\nAnswer: {text}\n\nVerdict (1 or 0):"
|
||||
)
|
||||
try:
|
||||
resp = await asyncio.wait_for(
|
||||
litellm.acompletion(
|
||||
model=JUDGE_MODEL,
|
||||
messages=[{"role": "user", "content": judge_prompt}],
|
||||
max_tokens=32,
|
||||
temperature=0.0,
|
||||
),
|
||||
timeout=JUDGE_TIMEOUT,
|
||||
)
|
||||
verdict = resp.choices[0].message.content.strip()
|
||||
update_progress(_session_idx, "judging")
|
||||
return 1 if verdict.startswith("1") else 0
|
||||
except Exception:
|
||||
update_progress(_session_idx, "judging")
|
||||
return -1
|
||||
|
||||
scores = await asyncio.gather(*[
|
||||
judge_one(m["prompt"], r["text"]) for m, r in zip(batch, responses)
|
||||
])
|
||||
|
||||
session_rewards: list[float] = []
|
||||
model_counts: dict[str, int] = {}
|
||||
for meta, dec, resp, score in zip(batch, decisions, responses, scores):
|
||||
model = dec["model"]
|
||||
tier = dec["tier"]
|
||||
model_counts[model] = model_counts.get(model, 0) + 1
|
||||
|
||||
if resp["success"] and score >= 0:
|
||||
reward = hybrid.record_quality(
|
||||
tier, model, resp["latency"], resp["tokens"],
|
||||
quality_score=float(score),
|
||||
quality_weight=QUALITY_WEIGHT,
|
||||
)
|
||||
elif resp["success"] and score == -1:
|
||||
reward = hybrid.record_success(tier, model, resp["latency"], resp["tokens"])
|
||||
else:
|
||||
hybrid.record_failure(tier, model)
|
||||
reward = 0.0
|
||||
|
||||
session_rewards.append(reward)
|
||||
all_results.append({
|
||||
"session": session_idx,
|
||||
"model": model,
|
||||
"tier": tier.value if hasattr(tier, "value") else tier,
|
||||
"category": meta["category"],
|
||||
"difficulty": meta["tier"],
|
||||
"latency": resp["latency"],
|
||||
"tokens": resp["tokens"],
|
||||
"success": resp["success"],
|
||||
"judge_score": score,
|
||||
"reward": reward,
|
||||
})
|
||||
|
||||
successful = [r for r in responses if r["success"]]
|
||||
avg_lat = sum(r["latency"] for r in successful) / len(successful) if successful else 0
|
||||
judged = [s for s in scores if s >= 0]
|
||||
accuracy = sum(1 for s in judged if s == 1) / len(judged) if judged else 0
|
||||
avg_reward = sum(session_rewards) / len(session_rewards) if session_rewards else 0
|
||||
|
||||
summary = {
|
||||
"session": session_idx,
|
||||
"n": len(batch),
|
||||
"model_dist": model_counts,
|
||||
"avg_reward": avg_reward,
|
||||
"avg_latency": avg_lat,
|
||||
"accuracy": accuracy,
|
||||
"wall_time": session_wall,
|
||||
}
|
||||
session_summaries.append(summary)
|
||||
|
||||
print(f"\n Session {session_idx+1}/{num_sessions}: "
|
||||
f"models={model_counts}, "
|
||||
f"acc={accuracy:.3f}, "
|
||||
f"lat={avg_lat:.3f}s, "
|
||||
f"reward={avg_reward:.3f}, "
|
||||
f"wall={session_wall:.1f}s")
|
||||
|
||||
total_wall = time.perf_counter() - t_total_start
|
||||
print(f"\r [{'#' * bar_width}] 100.0% | Done{' ' * 30}")
|
||||
print()
|
||||
|
||||
print(f"\n{'='*70}")
|
||||
print("FINAL RESULTS -- Hybrid Router v1.0")
|
||||
print(f"{'='*70}")
|
||||
|
||||
total_success = [r for r in all_results if r["success"]]
|
||||
total_judged = [r for r in all_results if r["judge_score"] >= 0]
|
||||
total_correct = [r for r in total_judged if r["judge_score"] == 1]
|
||||
|
||||
print(f"\n Total requests: {len(all_results)}")
|
||||
print(f" Success rate: {len(total_success)}/{len(all_results)}")
|
||||
print(f" Total wall time: {total_wall:.1f}s")
|
||||
if total_judged:
|
||||
print(f" Overall accuracy: {len(total_correct)}/{len(total_judged)} "
|
||||
f"({len(total_correct)/len(total_judged)*100:.1f}%)")
|
||||
overall_avg_lat = 0.0
|
||||
overall_tps = 0.0
|
||||
overall_total_tok = 0
|
||||
if total_success:
|
||||
overall_avg_lat = sum(r["latency"] for r in total_success) / len(total_success)
|
||||
overall_total_tok = sum(r["tokens"] for r in total_success)
|
||||
total_lat_sum = sum(r["latency"] for r in total_success)
|
||||
overall_tps = overall_total_tok / total_lat_sum if total_lat_sum > 0 else 0
|
||||
print(f" Avg latency: {overall_avg_lat:.3f}s")
|
||||
print(f" Total tokens: {overall_total_tok}")
|
||||
print(f" Throughput: {overall_tps:.1f} tok/s")
|
||||
|
||||
print(f"\n--- Per-Model Stats ---")
|
||||
models = sorted(set(r["model"] for r in all_results))
|
||||
for m in models:
|
||||
m_results = [r for r in all_results if r["model"] == m]
|
||||
m_success = [r for r in m_results if r["success"]]
|
||||
m_judged = [r for r in m_results if r["judge_score"] >= 0]
|
||||
m_correct = [r for r in m_judged if r["judge_score"] == 1]
|
||||
m_lat = sum(r["latency"] for r in m_success) / len(m_success) if m_success else 0
|
||||
m_lat_sum = sum(r["latency"] for r in m_success)
|
||||
m_tps = sum(r["tokens"] for r in m_success) / m_lat_sum if m_lat_sum > 0 else 0
|
||||
m_acc = len(m_correct) / len(m_judged) if m_judged else 0
|
||||
print(f" {m}: n={len(m_results)}, acc={m_acc:.3f}, avg_lat={m_lat:.3f}s, tps={m_tps:.1f}")
|
||||
|
||||
print(f"\n--- Per-Tier Stats ---")
|
||||
print(f" {'Tier':<10} {'N':>5} {'Acc':>6} {'AvgLat':>8} {'TPS':>7} {'Model Dist'}")
|
||||
for tier in ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]:
|
||||
tier_results = [r for r in all_results if r["tier"] == tier]
|
||||
if not tier_results:
|
||||
continue
|
||||
t_success = [r for r in tier_results if r["success"]]
|
||||
t_judged = [r for r in tier_results if r["judge_score"] >= 0]
|
||||
t_correct = [r for r in t_judged if r["judge_score"] == 1]
|
||||
t_acc = len(t_correct) / len(t_judged) if t_judged else 0
|
||||
t_lat = sum(r["latency"] for r in t_success) / len(t_success) if t_success else 0
|
||||
t_lat_sum = sum(r["latency"] for r in t_success)
|
||||
t_tps = sum(r["tokens"] for r in t_success) / t_lat_sum if t_lat_sum > 0 else 0
|
||||
model_dist: dict[str, int] = {}
|
||||
for r in tier_results:
|
||||
model_dist[r["model"]] = model_dist.get(r["model"], 0) + 1
|
||||
print(f" {tier:<10} {len(tier_results):>5} {t_acc:>5.3f} {t_lat:>8.3f} {t_tps:>6.1f} {model_dist}")
|
||||
|
||||
print(f"\n--- Per-Category Per-Difficulty ---")
|
||||
print(f" {'Cat/Diff':<12} {'N':>4} {'Acc':>6} {'AvgLat':>8} {'Model Dist'}")
|
||||
for cat in ("math", "code"):
|
||||
for pt in range(1, 6):
|
||||
cat_results = [r for r in all_results if r["category"] == cat and r["difficulty"] == pt]
|
||||
if not cat_results:
|
||||
continue
|
||||
c_success = [r for r in cat_results if r["success"]]
|
||||
c_judged = [r for r in cat_results if r["judge_score"] >= 0]
|
||||
c_correct = [r for r in c_judged if r["judge_score"] == 1]
|
||||
c_acc = len(c_correct) / len(c_judged) if c_judged else 0
|
||||
c_lat = sum(r["latency"] for r in c_success) / len(c_success) if c_success else 0
|
||||
model_dist = {}
|
||||
for r in cat_results:
|
||||
model_dist[r["model"]] = model_dist.get(r["model"], 0) + 1
|
||||
print(f" {cat}/t{pt:<8} {len(cat_results):>4} {c_acc:>5.3f} {c_lat:>8.3f} {model_dist}")
|
||||
|
||||
if not RUN_TIER_STATIC:
|
||||
print(f"\n--- Session Convergence ---")
|
||||
print(f" {'Session':<8} {'4B%':>6} {'9B%':>6} {'Acc':>6} {'Lat':>7} {'Reward':>7}")
|
||||
for s in session_summaries:
|
||||
n = s["n"]
|
||||
pct_4b = s["model_dist"].get("qwen-local-4b", 0) / n * 100
|
||||
pct_9b = s["model_dist"].get("qwen-local-9b", 0) / n * 100
|
||||
print(f" {s['session']+1:<8} {pct_4b:>5.1f}% {pct_9b:>5.1f}% "
|
||||
f"{s['accuracy']:>5.3f} {s['avg_latency']:>6.3f} {s['avg_reward']:>6.3f}")
|
||||
|
||||
print(f"\n--- Final Bandit State (Thompson posteriors) ---")
|
||||
state = hybrid.state()
|
||||
for tier, bandit_state in state["tier_bandits"].items():
|
||||
print(f" {tier}:")
|
||||
print(f" counts: {bandit_state['counts']}")
|
||||
means_str = {k: f"{v:.4f}" for k, v in bandit_state["means"].items()}
|
||||
print(f" means: {means_str}")
|
||||
if "alpha" in bandit_state:
|
||||
alpha_str = {k: f"{v:.2f}" for k, v in bandit_state["alpha"].items()}
|
||||
beta_str = {k: f"{v:.2f}" for k, v in bandit_state["beta"].items()}
|
||||
print(f" alpha: {alpha_str}")
|
||||
print(f" beta: {beta_str}")
|
||||
|
||||
rr_results: list[dict] = []
|
||||
if RUN_ROUND_ROBIN:
|
||||
print(f"\n{'='*70}")
|
||||
print("BASELINE -- Round-Robin (no intelligence)")
|
||||
print(f"{'='*70}")
|
||||
|
||||
rr_models = ["qwen-local-4b", "qwen-local-9b"]
|
||||
t_rr_start = time.perf_counter()
|
||||
|
||||
for session_idx in range(num_sessions):
|
||||
batch_start = session_idx * BATCH_SIZE
|
||||
batch_end = min(batch_start + BATCH_SIZE, len(prompts))
|
||||
batch = prompts[batch_start:batch_end]
|
||||
|
||||
async def rr_run_one(idx, meta, _session_idx=session_idx):
|
||||
model = rr_models[idx % len(rr_models)]
|
||||
async with req_semaphore:
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
resp = await router.acompletion(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": meta["prompt"]}],
|
||||
max_tokens=max_tokens_by_tier[meta["tier"]],
|
||||
)
|
||||
latency = time.perf_counter() - t0
|
||||
tokens = resp.usage.completion_tokens if resp.usage else 0
|
||||
text = resp.choices[0].message.content if resp.choices else ""
|
||||
return {"model": model, "latency": latency, "tokens": tokens, "text": text or "", "success": True}
|
||||
except Exception:
|
||||
return {"model": model, "latency": time.perf_counter() - t0, "tokens": 0, "text": "", "success": False}
|
||||
|
||||
global_offset = batch_start
|
||||
if CPU_INFERENCE:
|
||||
indexed_batch = list(enumerate(batch))
|
||||
responses = [None] * len(indexed_batch)
|
||||
model_assignments = [rr_models[(global_offset + j) % len(rr_models)] for j in range(len(batch))]
|
||||
models_in_batch = sorted(set(model_assignments))
|
||||
for model in models_in_batch:
|
||||
model_items = [(j, meta) for j, meta in indexed_batch if model_assignments[j] == model]
|
||||
model_responses = await asyncio.gather(*[
|
||||
rr_run_one(global_offset + j, meta) for j, meta in model_items
|
||||
])
|
||||
for (j, _), resp in zip(model_items, model_responses):
|
||||
responses[j] = resp
|
||||
else:
|
||||
responses = await asyncio.gather(*[
|
||||
rr_run_one(global_offset + j, meta) for j, meta in enumerate(batch)
|
||||
])
|
||||
|
||||
if skip_judge:
|
||||
scores = [1] * len(batch)
|
||||
else:
|
||||
async def rr_judge_one(prompt, text):
|
||||
if not text:
|
||||
return 0
|
||||
async with judge_semaphore:
|
||||
judge_prompt = (
|
||||
"You are an accuracy judge. Given a question and an answer, "
|
||||
"determine if the answer is correct.\n\n"
|
||||
"Respond with ONLY a single digit: 1 if the answer is correct, "
|
||||
"0 if it is incorrect or incomplete.\n\n"
|
||||
f"Question: {prompt}\n\nAnswer: {text}\n\nVerdict (1 or 0):"
|
||||
)
|
||||
try:
|
||||
resp = await asyncio.wait_for(
|
||||
litellm.acompletion(
|
||||
model=JUDGE_MODEL,
|
||||
messages=[{"role": "user", "content": judge_prompt}],
|
||||
max_tokens=32,
|
||||
temperature=0.0,
|
||||
),
|
||||
timeout=JUDGE_TIMEOUT,
|
||||
)
|
||||
verdict = resp.choices[0].message.content.strip()
|
||||
return 1 if verdict.startswith("1") else 0
|
||||
except Exception:
|
||||
return -1
|
||||
|
||||
scores = await asyncio.gather(*[
|
||||
rr_judge_one(m["prompt"], r["text"]) for m, r in zip(batch, responses)
|
||||
])
|
||||
|
||||
for meta, resp, score in zip(batch, responses, scores):
|
||||
rr_results.append({
|
||||
"session": session_idx,
|
||||
"model": resp["model"],
|
||||
"category": meta["category"],
|
||||
"difficulty": meta["tier"],
|
||||
"latency": resp["latency"],
|
||||
"tokens": resp["tokens"],
|
||||
"success": resp["success"],
|
||||
"judge_score": score,
|
||||
})
|
||||
|
||||
successful = [r for r in responses if r["success"]]
|
||||
avg_lat = sum(r["latency"] for r in successful) / len(successful) if successful else 0
|
||||
judged = [s for s in scores if s >= 0]
|
||||
accuracy = sum(1 for s in judged if s == 1) / len(judged) if judged else 0
|
||||
model_counts = {}
|
||||
for r in responses:
|
||||
model_counts[r["model"]] = model_counts.get(r["model"], 0) + 1
|
||||
print(f" Session {session_idx+1}/{num_sessions}: "
|
||||
f"models={model_counts}, acc={accuracy:.3f}, lat={avg_lat:.3f}s")
|
||||
|
||||
rr_wall = time.perf_counter() - t_rr_start
|
||||
rr_success = [r for r in rr_results if r["success"]]
|
||||
rr_judged = [r for r in rr_results if r["judge_score"] >= 0]
|
||||
rr_correct = [r for r in rr_judged if r["judge_score"] == 1]
|
||||
rr_avg_lat = sum(r["latency"] for r in rr_success) / len(rr_success) if rr_success else 0
|
||||
rr_total_tok = sum(r["tokens"] for r in rr_success)
|
||||
rr_lat_sum = sum(r["latency"] for r in rr_success)
|
||||
rr_tps = rr_total_tok / rr_lat_sum if rr_lat_sum > 0 else 0
|
||||
rr_acc = len(rr_correct) / len(rr_judged) if rr_judged else 0
|
||||
|
||||
print(f"\n Round-Robin Summary:")
|
||||
print(f" Requests: {len(rr_results)}, Success: {len(rr_success)}")
|
||||
print(f" Accuracy: {rr_acc:.3f} ({len(rr_correct)}/{len(rr_judged)})")
|
||||
print(f" Avg latency: {rr_avg_lat:.3f}s")
|
||||
print(f" Throughput: {rr_tps:.1f} tok/s")
|
||||
print(f" Wall time: {rr_wall:.1f}s")
|
||||
|
||||
if total_success:
|
||||
print(f"\n --- Comparison: Hybrid Router vs Round-Robin ---")
|
||||
print(f" {'Metric':<16} {'Hybrid':>10} {'RoundRobin':>12} {'Delta':>10}")
|
||||
hybrid_acc = len(total_correct) / len(total_judged) if total_judged else 0
|
||||
print(f" {'Accuracy':<16} {hybrid_acc:>9.3f} {rr_acc:>11.3f} {hybrid_acc - rr_acc:>+9.3f}")
|
||||
print(f" {'Avg Latency':<16} {overall_avg_lat:>9.3f}s {rr_avg_lat:>10.3f}s {overall_avg_lat - rr_avg_lat:>+9.3f}s")
|
||||
print(f" {'Throughput':<16} {overall_tps:>8.1f} t/s {rr_tps:>9.1f} t/s {overall_tps - rr_tps:>+8.1f} t/s")
|
||||
|
||||
mode = "tier-static" if RUN_TIER_STATIC else ("round-robin" if RUN_ROUND_ROBIN and not all_results else "bandit")
|
||||
output_path = "hybrid_v1_eval_results.json"
|
||||
output = {
|
||||
"config": {
|
||||
"mode": mode,
|
||||
"quality_weight": QUALITY_WEIGHT,
|
||||
"batch_size": BATCH_SIZE,
|
||||
"concurrency": CONCURRENCY,
|
||||
"tier_candidates": {k: list(v) for k, v in hybrid_config.tier_candidates.items()},
|
||||
"target_tpt": hybrid_config.target_tpt,
|
||||
"judge_model": JUDGE_MODEL,
|
||||
},
|
||||
"summary": {
|
||||
"total_requests": len(all_results),
|
||||
"total_success": len(total_success),
|
||||
"total_judged": len(total_judged),
|
||||
"total_correct": len(total_correct),
|
||||
"overall_accuracy": len(total_correct) / len(total_judged) if total_judged else 0,
|
||||
"avg_latency": overall_avg_lat,
|
||||
"throughput_tps": overall_tps,
|
||||
"total_wall_time": total_wall,
|
||||
},
|
||||
"session_summaries": session_summaries,
|
||||
"all_results": all_results,
|
||||
}
|
||||
if not RUN_TIER_STATIC:
|
||||
output["bandit_state"] = hybrid.state()
|
||||
if rr_results:
|
||||
output["round_robin_baseline"] = {
|
||||
"results": rr_results,
|
||||
"summary": {
|
||||
"total_requests": len(rr_results),
|
||||
"total_success": len(rr_success),
|
||||
"accuracy": rr_acc,
|
||||
"avg_latency": rr_avg_lat,
|
||||
"throughput_tps": rr_tps,
|
||||
"wall_time": rr_wall,
|
||||
},
|
||||
}
|
||||
with open(output_path, "w") as f:
|
||||
json.dump(output, f, indent=2)
|
||||
print(f"\n Results saved to {output_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
377
tests/test_litellm/router_strategy/test_hybrid_router.py
Normal file
377
tests/test_litellm/router_strategy/test_hybrid_router.py
Normal file
|
|
@ -0,0 +1,377 @@
|
|||
import random
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.router_strategy.complexity_router.config import ComplexityTier
|
||||
from litellm.router_strategy.hybrid_router import HybridRouter, HybridRouterConfig
|
||||
from litellm.router_strategy.hybrid_router.bandit import (
|
||||
EpsilonGreedy,
|
||||
ThompsonSampling,
|
||||
UCB,
|
||||
make_bandit,
|
||||
)
|
||||
|
||||
|
||||
TIER_CANDIDATES = {
|
||||
"SIMPLE": ("model-small", "model-medium"),
|
||||
"MEDIUM": ("model-medium", "model-large"),
|
||||
"COMPLEX": ("model-large",),
|
||||
"REASONING": ("model-large",),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config() -> HybridRouterConfig:
|
||||
return HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="thompson")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def router(config: HybridRouterConfig) -> HybridRouter:
|
||||
return HybridRouter(config)
|
||||
|
||||
|
||||
class TestMakeBandit:
|
||||
def test_thompson(self):
|
||||
bandit = make_bandit("thompson", ("a", "b"))
|
||||
assert isinstance(bandit, ThompsonSampling)
|
||||
|
||||
def test_ucb(self):
|
||||
bandit = make_bandit("ucb", ("a", "b"))
|
||||
assert isinstance(bandit, UCB)
|
||||
|
||||
def test_epsilon_greedy(self):
|
||||
bandit = make_bandit("epsilon_greedy", ("a", "b"))
|
||||
assert isinstance(bandit, EpsilonGreedy)
|
||||
|
||||
def test_unknown_raises(self):
|
||||
with pytest.raises(ValueError, match="Unknown bandit"):
|
||||
make_bandit("nonexistent", ("a", "b"))
|
||||
|
||||
def test_single_arm_raises(self):
|
||||
with pytest.raises(ValueError, match="at least 2 arms"):
|
||||
make_bandit("thompson", ("only-one",))
|
||||
|
||||
|
||||
class TestUCB:
|
||||
def test_unexplored_arm_picked_first(self):
|
||||
bandit = UCB(("a", "b"), delta=0.1)
|
||||
first = bandit.select_arm()
|
||||
assert first in ("a", "b")
|
||||
|
||||
def test_explored_arms_use_ucb_score(self):
|
||||
bandit = UCB(("a", "b"), delta=0.1)
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
assert bandit.select_arm() == "a"
|
||||
|
||||
def test_select_arm_from_respects_eligible(self):
|
||||
bandit = UCB(("a", "b", "c"), delta=0.1)
|
||||
for _ in range(10):
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
bandit.update("c", 0.5)
|
||||
picked = bandit.select_arm_from(("b", "c"))
|
||||
assert picked in ("b", "c")
|
||||
|
||||
def test_invalid_delta_raises(self):
|
||||
with pytest.raises(ValueError, match="delta must be in"):
|
||||
UCB(("a", "b"), delta=0.0)
|
||||
|
||||
def test_state_tracks_counts(self):
|
||||
bandit = UCB(("a", "b"), delta=0.1)
|
||||
bandit.update("a", 0.8)
|
||||
bandit.update("a", 0.6)
|
||||
bandit.update("b", 0.4)
|
||||
state = bandit.state()
|
||||
assert state["counts"]["a"] == 2
|
||||
assert state["counts"]["b"] == 1
|
||||
assert abs(state["means"]["a"] - 0.7) < 1e-9
|
||||
|
||||
|
||||
class TestEpsilonGreedy:
|
||||
def test_unexplored_arm_picked_first(self):
|
||||
bandit = EpsilonGreedy(("a", "b"), epsilon=0.1)
|
||||
first = bandit.select_arm()
|
||||
assert first in ("a", "b")
|
||||
|
||||
def test_greedy_picks_best_after_exploration(self):
|
||||
bandit = EpsilonGreedy(("a", "b"), epsilon=0.0)
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
assert bandit.select_arm() == "a"
|
||||
|
||||
def test_full_exploration_picks_randomly(self):
|
||||
random.seed(42)
|
||||
bandit = EpsilonGreedy(("a", "b"), epsilon=1.0)
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
picks = {bandit.select_arm() for _ in range(50)}
|
||||
assert picks == {"a", "b"}
|
||||
|
||||
def test_invalid_epsilon_raises(self):
|
||||
with pytest.raises(ValueError, match="epsilon must be in"):
|
||||
EpsilonGreedy(("a", "b"), epsilon=1.5)
|
||||
|
||||
def test_state_tracks_counts(self):
|
||||
bandit = EpsilonGreedy(("a", "b"), epsilon=0.1)
|
||||
bandit.update("a", 0.5)
|
||||
bandit.update("b", 0.3)
|
||||
state = bandit.state()
|
||||
assert state["counts"]["a"] == 1
|
||||
assert state["counts"]["b"] == 1
|
||||
assert state["algorithm"] == "EpsilonGreedy"
|
||||
|
||||
|
||||
class TestThompsonSampling:
|
||||
def test_selects_from_arms(self):
|
||||
bandit = ThompsonSampling(("a", "b"))
|
||||
assert bandit.select_arm() in ("a", "b")
|
||||
|
||||
def test_update_shifts_posterior(self):
|
||||
bandit = ThompsonSampling(("a", "b"))
|
||||
for _ in range(20):
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
state = bandit.state()
|
||||
assert state["means"]["a"] > state["means"]["b"]
|
||||
|
||||
def test_custom_priors(self):
|
||||
bandit = ThompsonSampling(("a", "b"), arm_priors={"a": (10.0, 1.0), "b": (1.0, 10.0)})
|
||||
state = bandit.state()
|
||||
assert state["alpha"]["a"] == 10.0
|
||||
assert state["beta"]["b"] == 10.0
|
||||
|
||||
def test_select_arm_from_respects_eligible(self):
|
||||
bandit = ThompsonSampling(("a", "b", "c"))
|
||||
for _ in range(20):
|
||||
bandit.update("a", 1.0)
|
||||
bandit.update("b", 0.0)
|
||||
bandit.update("c", 0.5)
|
||||
picked = bandit.select_arm_from(("b", "c"))
|
||||
assert picked in ("b", "c")
|
||||
|
||||
def test_state_reports_algorithm(self):
|
||||
bandit = ThompsonSampling(("a", "b"))
|
||||
assert bandit.state()["algorithm"] == "ThompsonSampling"
|
||||
|
||||
|
||||
class TestClassification:
|
||||
def test_simple_greeting(self, router: HybridRouter):
|
||||
tier, score, signals = router.classify("hello")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
def test_code_request(self, router: HybridRouter):
|
||||
tier, score, signals = router.classify(
|
||||
"Write a python function to implement a distributed database query with async error handling"
|
||||
)
|
||||
assert tier in (ComplexityTier.MEDIUM, ComplexityTier.COMPLEX, ComplexityTier.REASONING)
|
||||
assert score > 0
|
||||
|
||||
def test_reasoning_keywords_trigger_reasoning_tier(self, router: HybridRouter):
|
||||
tier, score, signals = router.classify(
|
||||
"Think through step by step and reason through the chain of thought for this problem"
|
||||
)
|
||||
assert tier == ComplexityTier.REASONING
|
||||
|
||||
def test_short_prompt_gets_simple_tier(self, router: HybridRouter):
|
||||
tier, score, signals = router.classify("hi")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
def test_long_prompt_boosts_score(self, router: HybridRouter):
|
||||
long_prompt = "Explain " + "the architecture of distributed systems " * 50
|
||||
tier, score, _ = router.classify(long_prompt)
|
||||
assert score > 0
|
||||
|
||||
def test_multi_step_pattern_detected(self, router: HybridRouter):
|
||||
tier, score, _ = router.classify("First do X, then do Y. Step 1: read the file. Step 2: parse it.")
|
||||
assert score > 0
|
||||
|
||||
def test_simple_keywords_reduce_score(self, router: HybridRouter):
|
||||
tier, score, _ = router.classify("What is the definition of hello?")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
|
||||
|
||||
class TestGetCandidates:
|
||||
def test_returns_tier_candidates(self, router: HybridRouter):
|
||||
candidates = router.get_candidates(ComplexityTier.SIMPLE)
|
||||
assert candidates == ("model-small", "model-medium")
|
||||
|
||||
def test_returns_complex_candidates(self, router: HybridRouter):
|
||||
candidates = router.get_candidates(ComplexityTier.COMPLEX)
|
||||
assert candidates == ("model-large",)
|
||||
|
||||
def test_fallback_when_tier_missing(self):
|
||||
config = HybridRouterConfig(
|
||||
tier_candidates={"SIMPLE": ("model-small",)},
|
||||
bandit="thompson",
|
||||
)
|
||||
router = HybridRouter(config)
|
||||
candidates = router.get_candidates(ComplexityTier.COMPLEX)
|
||||
assert len(candidates) > 0
|
||||
|
||||
|
||||
class TestPickModel:
|
||||
def test_single_candidate_returns_it(self, router: HybridRouter):
|
||||
model = router.pick_model(ComplexityTier.COMPLEX)
|
||||
assert model == "model-large"
|
||||
|
||||
def test_multi_candidate_returns_valid_model(self, router: HybridRouter):
|
||||
model = router.pick_model(ComplexityTier.SIMPLE)
|
||||
assert model in ("model-small", "model-medium")
|
||||
|
||||
|
||||
class TestRoute:
|
||||
def test_returns_tier_and_model(self, router: HybridRouter):
|
||||
tier, model = router.route("hello")
|
||||
assert isinstance(tier, ComplexityTier)
|
||||
assert model in ("model-small", "model-medium", "model-large")
|
||||
|
||||
def test_simple_routes_to_simple_candidates(self, router: HybridRouter):
|
||||
tier, model = router.route("hi")
|
||||
assert tier == ComplexityTier.SIMPLE
|
||||
assert model in ("model-small", "model-medium")
|
||||
|
||||
|
||||
class TestRewardRecording:
|
||||
def test_record_success_returns_reward_in_range(self, router: HybridRouter):
|
||||
reward = router.record_success(ComplexityTier.SIMPLE, "model-small", latency_seconds=0.5, completion_tokens=10)
|
||||
assert 0.0 < reward <= 1.0
|
||||
|
||||
def test_record_success_updates_bandit_state(self, router: HybridRouter):
|
||||
state_before = router.state()
|
||||
count_before = state_before["tier_bandits"]["SIMPLE"]["counts"]["model-small"]
|
||||
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-small", latency_seconds=0.5, completion_tokens=10)
|
||||
|
||||
state_after = router.state()
|
||||
count_after = state_after["tier_bandits"]["SIMPLE"]["counts"]["model-small"]
|
||||
assert count_after == count_before + 1
|
||||
|
||||
def test_record_success_no_bandit_for_single_candidate_tier(self, router: HybridRouter):
|
||||
reward = router.record_success(ComplexityTier.COMPLEX, "model-large", latency_seconds=0.5, completion_tokens=10)
|
||||
assert 0.0 < reward <= 1.0
|
||||
assert "COMPLEX" not in router.state()["tier_bandits"]
|
||||
|
||||
def test_record_failure_gives_zero_reward(self, router: HybridRouter):
|
||||
router.record_failure(ComplexityTier.SIMPLE, "model-small")
|
||||
state = router.state()
|
||||
assert state["tier_bandits"]["SIMPLE"]["counts"]["model-small"] == 1
|
||||
|
||||
def test_record_quality_blends_efficiency_and_quality(self, router: HybridRouter):
|
||||
reward_high_quality = router.record_quality(
|
||||
ComplexityTier.SIMPLE,
|
||||
"model-small",
|
||||
latency_seconds=0.5,
|
||||
completion_tokens=10,
|
||||
quality_score=1.0,
|
||||
quality_weight=0.5,
|
||||
)
|
||||
router2 = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="thompson"))
|
||||
reward_low_quality = router2.record_quality(
|
||||
ComplexityTier.SIMPLE,
|
||||
"model-small",
|
||||
latency_seconds=0.5,
|
||||
completion_tokens=10,
|
||||
quality_score=0.0,
|
||||
quality_weight=0.5,
|
||||
)
|
||||
assert reward_high_quality > reward_low_quality
|
||||
|
||||
def test_record_quality_weight_zero_equals_efficiency_only(self, router: HybridRouter):
|
||||
reward_eff = router.record_success(
|
||||
ComplexityTier.SIMPLE, "model-small", latency_seconds=0.5, completion_tokens=10
|
||||
)
|
||||
router2 = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="thompson"))
|
||||
reward_qw0 = router2.record_quality(
|
||||
ComplexityTier.SIMPLE,
|
||||
"model-small",
|
||||
latency_seconds=0.5,
|
||||
completion_tokens=10,
|
||||
quality_score=0.0,
|
||||
quality_weight=0.0,
|
||||
)
|
||||
assert abs(reward_eff - reward_qw0) < 1e-9
|
||||
|
||||
def test_record_quality_weight_one_equals_quality_score(self):
|
||||
router = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="thompson"))
|
||||
reward = router.record_quality(
|
||||
ComplexityTier.SIMPLE,
|
||||
"model-small",
|
||||
latency_seconds=100.0,
|
||||
completion_tokens=1,
|
||||
quality_score=0.9,
|
||||
quality_weight=1.0,
|
||||
)
|
||||
assert abs(reward - 0.9) < 1e-9
|
||||
|
||||
|
||||
class TestBanditConvergence:
|
||||
def test_thompson_converges_to_faster_arm(self):
|
||||
router = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="thompson"))
|
||||
for _ in range(100):
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-small", latency_seconds=0.1, completion_tokens=50)
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-medium", latency_seconds=1.0, completion_tokens=50)
|
||||
picks = [router.pick_model(ComplexityTier.SIMPLE) for _ in range(50)]
|
||||
assert picks.count("model-small") > 30
|
||||
|
||||
def test_ucb_converges_to_faster_arm(self):
|
||||
router = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="ucb"))
|
||||
for _ in range(100):
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-small", latency_seconds=0.1, completion_tokens=50)
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-medium", latency_seconds=1.0, completion_tokens=50)
|
||||
picks = [router.pick_model(ComplexityTier.SIMPLE) for _ in range(50)]
|
||||
assert picks.count("model-small") > 30
|
||||
|
||||
def test_epsilon_greedy_converges_to_faster_arm(self):
|
||||
router = HybridRouter(HybridRouterConfig(tier_candidates=TIER_CANDIDATES, bandit="epsilon_greedy"))
|
||||
for _ in range(100):
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-small", latency_seconds=0.1, completion_tokens=50)
|
||||
router.record_success(ComplexityTier.SIMPLE, "model-medium", latency_seconds=1.0, completion_tokens=50)
|
||||
picks = [router.pick_model(ComplexityTier.SIMPLE) for _ in range(50)]
|
||||
assert picks.count("model-small") > 30
|
||||
|
||||
|
||||
class TestTierPriors:
|
||||
def test_priors_bias_initial_selection(self):
|
||||
config = HybridRouterConfig(
|
||||
tier_candidates=TIER_CANDIDATES,
|
||||
bandit="thompson",
|
||||
tier_priors={"SIMPLE": {"model-small": (20.0, 1.0), "model-medium": (1.0, 20.0)}},
|
||||
)
|
||||
router = HybridRouter(config)
|
||||
picks = [router.pick_model(ComplexityTier.SIMPLE) for _ in range(50)]
|
||||
assert picks.count("model-small") > 40
|
||||
|
||||
|
||||
class TestState:
|
||||
def test_state_includes_config(self, router: HybridRouter):
|
||||
state = router.state()
|
||||
assert "config" in state
|
||||
assert state["config"]["bandit"] == "thompson"
|
||||
|
||||
def test_state_includes_tier_bandits(self, router: HybridRouter):
|
||||
state = router.state()
|
||||
assert "tier_bandits" in state
|
||||
assert "SIMPLE" in state["tier_bandits"]
|
||||
assert "MEDIUM" in state["tier_bandits"]
|
||||
assert "COMPLEX" not in state["tier_bandits"]
|
||||
|
||||
|
||||
class TestHybridRouterConfig:
|
||||
def test_frozen_config(self):
|
||||
config = HybridRouterConfig(tier_candidates=TIER_CANDIDATES)
|
||||
with pytest.raises(AttributeError):
|
||||
config.bandit = "ucb"
|
||||
|
||||
def test_default_bandit_is_thompson(self):
|
||||
config = HybridRouterConfig(tier_candidates=TIER_CANDIDATES)
|
||||
assert config.bandit == "thompson"
|
||||
|
||||
def test_custom_keywords_override_defaults(self):
|
||||
config = HybridRouterConfig(tier_candidates=TIER_CANDIDATES, code_keywords=("custom_kw",))
|
||||
router = HybridRouter(config)
|
||||
assert router._code_keywords == ("custom_kw",)
|
||||
Loading…
Add table
Reference in a new issue