mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
Merge 3ae085c639 into 2dccc0dc79
This commit is contained in:
commit
574f5c0042
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