This commit is contained in:
souvikku 2026-09-23 14:46:15 +00:00 • committed by GitHub
commit 574f5c0042
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 2849 additions and 0 deletions

View file

@ -0,0 +1,6 @@
from litellm.router_strategy.hybrid_router.hybrid_router import (
HybridRouter,
HybridRouterConfig,
)
__all__ = ["HybridRouter", "HybridRouterConfig"]

View 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)

View 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))

File diff suppressed because it is too large Load diff

View 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())

View 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",)