From 3ae085c6394e3c1b83191d3a8836f5a2fae48d82 Mon Sep 17 00:00:00 2001 From: "souvikk.kundu" Date: Wed, 2 Sep 2026 20:17:04 -0700 Subject: [PATCH] feat(router_strategy): add hybrid router combining complexity gating with per-tier MAB selection Two-layer routing strategy: Layer 1 classifies request complexity into SIMPLE/MEDIUM/COMPLEX/REASONING tiers using heuristic scoring, Layer 2 picks the best model within each tier using a per-tier multi-armed bandit (Thompson Sampling, UCB, or Epsilon-Greedy). Per-tier bandits learn independently, so a model that is fast for simple queries but slow for complex ones gets scored separately at each tier. Includes eval script and 50 unit tests covering all three bandit algorithms, classification, reward recording, convergence, and tier priors. --- .../router_strategy/hybrid_router/__init__.py | 6 + .../router_strategy/hybrid_router/bandit.py | 252 ++++ .../hybrid_router/hybrid_router.py | 318 +++++ scripts/efficiency_eval_prompts.py | 1179 +++++++++++++++++ scripts/test_custom_router_hybrid_v1.py | 717 ++++++++++ .../router_strategy/test_hybrid_router.py | 377 ++++++ 6 files changed, 2849 insertions(+) create mode 100644 litellm/router_strategy/hybrid_router/__init__.py create mode 100644 litellm/router_strategy/hybrid_router/bandit.py create mode 100644 litellm/router_strategy/hybrid_router/hybrid_router.py create mode 100644 scripts/efficiency_eval_prompts.py create mode 100644 scripts/test_custom_router_hybrid_v1.py create mode 100644 tests/test_litellm/router_strategy/test_hybrid_router.py diff --git a/litellm/router_strategy/hybrid_router/__init__.py b/litellm/router_strategy/hybrid_router/__init__.py new file mode 100644 index 00000000000..6ce813658c4 --- /dev/null +++ b/litellm/router_strategy/hybrid_router/__init__.py @@ -0,0 +1,6 @@ +from litellm.router_strategy.hybrid_router.hybrid_router import ( + HybridRouter, + HybridRouterConfig, +) + +__all__ = ["HybridRouter", "HybridRouterConfig"] diff --git a/litellm/router_strategy/hybrid_router/bandit.py b/litellm/router_strategy/hybrid_router/bandit.py new file mode 100644 index 00000000000..7fda18fc732 --- /dev/null +++ b/litellm/router_strategy/hybrid_router/bandit.py @@ -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) diff --git a/litellm/router_strategy/hybrid_router/hybrid_router.py b/litellm/router_strategy/hybrid_router/hybrid_router.py new file mode 100644 index 00000000000..21ac451004d --- /dev/null +++ b/litellm/router_strategy/hybrid_router/hybrid_router.py @@ -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)) diff --git a/scripts/efficiency_eval_prompts.py b/scripts/efficiency_eval_prompts.py new file mode 100644 index 00000000000..dffefe1c362 --- /dev/null +++ b/scripts/efficiency_eval_prompts.py @@ -0,0 +1,1179 @@ +""" +1000 evaluation prompts designed to reveal Thompson Sampling bandit advantage +over static tier-based routing. + +Key design principle: the keyword-based classifier's tier assignments DON'T +perfectly predict which model (4B vs 9B) will succeed. This creates room for +the bandit to learn better per-tier distributions than static routing provides. + +Structure: + - SIMPLE tier prompts (classified SIMPLE by keyword/length heuristics): + - 70% genuinely trivial (4B handles fine) -> static SIMPLE->4B is correct + - 30% "deceptively simple" (short/simple keywords but actually hard for 4B) + -> static loses here; bandit learns to shift toward 9B + + - MEDIUM tier prompts (classified MEDIUM due to code/technical keywords): + - 50% truly need 9B (complex logic despite moderate classification) + - 50% "easy mediums" (trigger keywords but actually trivial for 4B) + -> static wastes latency sending all MEDIUM->9B; bandit learns 4B is fine + +Two categories: math/logic (500) and code (500). +Within each, prompts are tagged with expected difficulty (1-5) independent +of how the classifier routes them. +""" + +# ─── MATH: SIMPLE-classified prompts ─── +# These will classify as SIMPLE (short, contain "what is", "how many", etc.) + +# Genuinely trivial (4B should get these right) +MATH_SIMPLE_EASY = [ + "What is 7 * 8?", + "What is 15% of 200?", + "What is the square root of 144?", + "What is 2^10?", + "How many seconds are in one hour?", + "What is 0.25 as a fraction?", + "What is 40% of 80?", + "What is the LCM of 4 and 6?", + "How many degrees in a right angle?", + "What is 3^4?", + "What is 25% of 400?", + "What is 10^6?", + "How many edges does a cube have?", + "What is 50% of 50?", + "What is log base 10 of 1000?", + "What is 12 * 12?", + "What is the perimeter of a rectangle 3 by 7?", + "What is 9 * 11?", + "What is the volume of a cube with side 3?", + "What is 1% of 5000?", + "How many zeros in one million?", + "What is 13^2?", + "How many sides does a decagon have?", + "What is 256 / 16?", + "What is 3/5 as a percentage?", + "What is log base 2 of 64?", + "What is 11 * 11?", + "How many days in a leap year?", + "What is 4!?", + "What is 999 - 111?", + "How many inches in a foot?", + "What is 7/8 as a decimal?", + "What is 20 * 0.5?", + "What is the sum of the first 5 positive integers?", + "What is the square of 15?", + "How many faces does an octahedron have?", + "What is 1000 * 0.01?", + "What is 45 + 55?", + "How many right angles in a rectangle?", + "What is 5% of 1000?", + "What is 8^2 - 6^2?", + "What is the product of 3 and 17?", + "How many months have 31 days?", + "What is 14 * 7?", + "What is the absolute value of -42?", + "What is 200 / 0.5?", + "How many quarts in a gallon?", + "What is 1.5 squared?", + "What is 100 - 37?", + "What is 2.5 * 4?", + "What is the next prime after 7?", + "What is |-5| + |3|?", + "What is the median of {1, 3, 5, 7, 9}?", + "What is the reciprocal of 5?", + "How many faces does a tetrahedron have?", + "What is 3.14 rounded to one decimal?", + "How many vertices does a pentagon have?", + "What is 1/3 + 1/6?", + "What is 2^0?", + "What is the next perfect square after 25?", + "What is the GCD of 12 and 18?", + "What is 5^3?", + "How many centimeters in 2.5 meters?", + "What is the area of a triangle with base 10 and height 6?", + "What is 10! / 9!?", + "What is 0.5^2?", + "How many prime numbers are less than 20?", + "What is the smallest 3-digit number?", + "What is the mean of {2, 4, 6, 8, 10}?", + "What is the circumference of a circle with diameter 10? Leave in terms of pi.", + # 70 easy prompts +] + +# Deceptively simple: short, contain simple keywords, but actually hard for 4B +MATH_SIMPLE_HARD = [ + "What is 0.1 + 0.2? Give the exact decimal.", + "Is 91 prime?", + "What is (-2)^(-3)?", + "What is the digit in the units place of 7^100?", + "What is 0.999... as a fraction?", + "How many trailing zeros does 100! have?", + "What is the remainder when 2^100 is divided by 7?", + "What is sqrt(-1) * sqrt(-1)?", + "How many factors does 360 have?", + "What is the sum of all integers from -50 to 50?", + "What is 1/7 as a repeating decimal?", + "What is the 100th digit of pi?", + "How many primes are between 100 and 120?", + "What is the last digit of 3^2023?", + "Is 2^31 - 1 prime?", + "What is the smallest number divisible by 1 through 10?", + "What is floor(-3.7)?", + "How many subsets does a set of 10 elements have?", + "What is the sum 1/1 + 1/2 + 1/3 + ... + 1/10 as a fraction?", + "What is 123456789 mod 9?", + "How many rectangles are on a standard 8x8 chessboard?", + "What is the maximum product of integers that sum to 20?", + "Is 561 a Carmichael number?", + "How many digits does 2^100 have?", + "What is the smallest prime larger than 1000?", + "What is (-1)^(1/3)?", + "What is the sum of the reciprocals of the first 5 primes?", + "How many ways can you make change for a dollar using quarters, dimes, nickels, and pennies?", + "What is the determinant of [[1,2],[3,4]]?", + "Is the number 2^67 - 1 prime?", +] + +# ─── MATH: MEDIUM-classified prompts ─── +# These trigger code/technical/reasoning keywords enough to classify as MEDIUM + +# Easy mediums: keywords trigger MEDIUM but 4B can solve them +MATH_MEDIUM_EASY = [ + "Implement the formula to convert Fahrenheit to Celsius and compute 72F.", + "Using the algorithm for GCD, find GCD(48, 36).", + "Optimize this calculation: what is 99 * 101 using the difference of squares?", + "Debug this: if x = 5, what is 2x + 3?", + "Implement a function that returns whether 15 is odd. What does it return?", + "Using the quadratic formula, what are the roots of x^2 - 5x + 6 = 0?", + "What is the performance difference between O(n) and O(n^2) when n=10?", + "Implement the algorithm: what is the 5th Fibonacci number?", + "Optimize: what is 25 * 4 * 7 * 0?", + "Debug: if a list is [1,2,3], what is its length?", + "Using the algorithm for binary search, how many steps to find 7 in [1,2,3,4,5,6,7,8]?", + "Implement the formula for the area of a circle with radius 5.", + "What is the throughput if 100 items are processed in 10 seconds?", + "Using the formula for compound interest, what is $100 at 10% for 1 year?", + "Implement: given a sorted list [1,3,5,7,9], is 5 in it?", + "What is the latency if a request travels 300km at the speed of light?", + "Optimize the expression: (a+b)^2 when a=3, b=4.", + "Debug this logic: if x > 5 and x < 3, is this ever true?", + "Implement the algorithm to reverse a 3-element list [a, b, c].", + "Using performance analysis, which is faster: 2*n or n+n?", + "What is the memory needed to store 1000 integers of 4 bytes each?", + "Implement: what is max(3, 7, 1, 9, 2)?", + "Using the algorithm for factorial, what is 5!?", + "What is the cpu time for 10^9 operations at 1 GHz?", + "Optimize: how do you multiply by 8 using bit shifts?", + "Debug: what is the output of 'print(2 + 2)' in Python?", + "Implement the formula for the slope between (0,0) and (3,6).", + "What is the throughput of a database processing 1000 queries/second?", + "Using the algorithm for selection sort on [3,1,2], what is the result?", + "Implement: is the string 'racecar' a palindrome?", + "What is the latency of a function that runs in O(1)?", + "Using parallel processing with 4 cores, how long does a 4-second serial task take ideally?", + "Optimize: what is faster, checking n items linearly or using a hash table?", + "Debug: what does range(5) produce in Python?", + "Implement the formula for the distance between (0,0) and (3,4).", + "What is the memory usage of an empty Python list?", + "Using the algorithm for bubble sort, sort [3,2,1].", + "What is the performance of looking up a key in a Python dict?", + "Implement: convert the binary number 1010 to decimal.", + "What is the cpu utilization if 2 of 4 cores are busy?", + "Using the algorithm for linear search, find 3 in [1,2,3,4,5].", + "Debug: what is bool(0) in Python?", + "Implement the formula for simple interest on $200 at 5% for 2 years.", + "What is the throughput in MB/s if 1GB transfers in 10 seconds?", + "Using optimization, simplify: x^2 - x^2 + 5.", + "Implement: what is the sum of even numbers from 1 to 10?", + "What is the latency added by 10 sequential function calls each taking 1ms?", + "Debug: what does len('hello world') return?", + "Using the algorithm for counting sort, sort [2,1,1,3,2].", + "Implement: what is the bitwise AND of 12 and 10?", + # 50 easy mediums +] + +# Hard mediums: correctly classified as MEDIUM, actually need 9B +MATH_MEDIUM_HARD = [ + "Implement an algorithm to determine whether the equation x^3 + y^3 = z^3 has solutions for positive integers less than 100.", + "Using the algorithm for matrix multiplication, compute [[1,2],[3,4]] * [[5,6],[7,8]].", + "Optimize the algorithm: find the maximum subarray sum of [-2, 1, -3, 4, -1, 2, 1, -5, 4].", + "Debug this algorithm: merge sort claims [3,1,4,1,5] sorts to [1,1,3,4,5]. Verify each merge step.", + "Implement the algorithm to find all prime factors of 2310.", + "Using the Euclidean algorithm extended version, find integers x,y such that 35x + 15y = 5.", + "Optimize the calculation: find the determinant of a 3x3 matrix [[2,1,3],[0,4,1],[5,2,1]].", + "Implement the algorithm for finding the longest increasing subsequence of [10,9,2,5,3,7,101,18].", + "Using the algorithm for topological sort, order the tasks: A->B, A->C, B->D, C->D.", + "Debug this dynamic programming solution: the correct number of ways to climb 5 stairs (1 or 2 steps) is what?", + "Implement the algorithm to solve the Tower of Hanoi for 4 disks. How many moves?", + "Using the algorithm for Dijkstra, find the shortest path in a graph with edges: A-B:4, A-C:2, B-D:5, C-B:1, C-D:8.", + "Optimize: find the minimum number of coins (1,5,10,25) to make 67 cents.", + "Implement the algorithm to check if the graph with edges {(A,B),(B,C),(C,A),(A,D)} has an Eulerian path.", + "Using the algorithm for computing convex hull, which points are vertices of the hull: (0,0),(1,1),(2,0),(1,3),(0,2)?", + "Debug: the knapsack problem with items [(w=2,v=3),(w=3,v=4),(w=4,v=5)] and capacity 5. What is the max value?", + "Implement the algorithm to convert the regular expression (a|b)*abb to a DFA. How many states?", + "Using numerical methods, approximate sqrt(2) using Newton's method starting from x=1 for 3 iterations.", + "Optimize: find the maximum flow in a network with edges s->a:10, s->b:5, a->b:15, a->t:5, b->t:10.", + "Implement the algorithm for computing edit distance between 'kitten' and 'sitting'.", + "Using the algorithm for fast exponentiation, compute 3^13 mod 7. Show the squaring steps.", + "Debug: Floyd-Warshall on a 4-node graph. Find all shortest paths.", + "Implement the algorithm for Huffman coding of frequencies: A=5, B=9, C=12, D=13, E=16, F=45.", + "Using the algorithm for computing eigenvalues, find the eigenvalues of [[4,-2],[1,1]].", + "Optimize: solve the rod cutting problem for a rod of length 8 with prices [1,5,8,9,10,17,17,20].", + "Implement the algorithm to find the strongly connected components of a graph with edges: 1->2, 2->3, 3->1, 2->4, 4->5.", + "Using the CRT algorithm, solve: x ≡ 2 (mod 3), x ≡ 3 (mod 5), x ≡ 2 (mod 7).", + "Debug: verify that the FFT of [1,1,1,1,0,0,0,0] gives the correct frequency components.", + "Implement the algorithm for Gaussian elimination on [[2,1,-1,8],[-3,-1,2,-11],[-2,1,2,-3]].", + "Using the algorithm for computing the permanent of a matrix, find perm([[1,2],[3,4]]).", + "Optimize the dynamic programming: find the longest common subsequence of 'ABCBDAB' and 'BDCAB'.", + "Implement the algorithm to determine the chromatic number of the Petersen graph.", + "Using the algorithm for singular value decomposition, find the rank of [[1,2,3],[4,5,6],[7,8,9]].", + "Debug: verify that Dijkstra's algorithm fails on negative edge weights. Give a counterexample.", + "Implement the algorithm for the stable matching problem (Gale-Shapley) with 3 men and 3 women.", + "Using the algorithm for polynomial interpolation, find the polynomial through (0,1),(1,2),(2,5).", + "Optimize: find the minimum spanning tree of a graph with edges A-B:4, A-C:8, B-C:2, B-D:5, C-D:3.", + "Implement the algorithm for computing the Catalan number C(10) using dynamic programming.", + "Using the algorithm for linear programming (simplex), maximize 3x+5y subject to x+y<=4, x<=3, y<=3.", + "Debug: verify the correctness of quicksort's partition step on [3,6,8,10,1,2,1].", + "Implement the algorithm for finding articulation points in the graph: 1-2, 2-3, 3-4, 4-2, 3-5.", + "Using the algorithm for matrix chain multiplication, find the optimal parenthesization for dimensions [10,30,5,60].", + "Optimize: find the minimum path sum in the grid [[1,3,1],[1,5,1],[4,2,1]] from top-left to bottom-right.", + "Implement the algorithm for the maximum bipartite matching in a graph with edges: a-1, a-2, b-2, b-3, c-1.", + "Using the algorithm for computing Fibonacci with matrix exponentiation, find F(20).", + "Debug: verify that the Bellman-Ford algorithm correctly detects the negative cycle in: A->B:1, B->C:2, C->A:-4.", + "Implement the algorithm for finding the minimum vertex cover of a bipartite graph.", + "Using the algorithm for polynomial multiplication via FFT, multiply (1+2x+3x^2) by (4+5x).", + "Optimize: solve the egg drop problem with 2 eggs and 100 floors. What is the minimum worst-case trials?", + "Implement the algorithm for computing the number of spanning trees of K5 using Kirchhoff's theorem.", + # 50 hard mediums +] + +# ─── MATH: Additional SIMPLE prompts to fill quota ─── +# More tricky SIMPLE-classified prompts (short, keyword-heavy, but challenging) +MATH_SIMPLE_TRICKY = [ + "What is the sum of all prime numbers less than 30?", + "How many perfect squares are between 1 and 1000?", + "What is 2^2^3?", + "Is 0 a natural number?", + "What is the 10th prime number?", + "How many divisors does 120 have?", + "What is the sum of the digits of 999999999999?", + "Is infinity a number?", + "What is 0^0?", + "How many triangles in a complete graph K6?", + "What is the largest prime factor of 1001?", + "Is sqrt(4) rational?", + "What is the GCD of 0 and 5?", + "How many zeros does the function x^3 - x have?", + "What is the sum 1 + 2 + 4 + 8 + ... + 2^10?", + "Is pi + e rational?", + "What is the 20th Fibonacci number?", + "How many diagonals does a 12-sided polygon have?", + "What is the smallest perfect number?", + "Is -0 equal to 0?", + "What is the cube root of -27?", + "How many 3-digit palindromes exist?", + "What is the product of all single-digit primes?", + "Is 2/0 infinity or undefined?", + "What is the angle between clock hands at 9:15?", + "How many positive integers less than 100 are coprime to 100?", + "What is the sum of interior angles of a 20-gon?", + "Is every prime number odd?", + "What is the digital root of 123456?", + "How many perfect cubes are less than 1000?", + "What is i^i where i is the imaginary unit?", + "Is the empty set a subset of every set?", + "What is the largest number you can make with three 4s and basic operations?", + "How many distinct handshakes in a group of 20 people?", + "What is the probability of rolling a sum of 7 with two dice?", + "Is 0.999... equal to 1?", + "What is the value of the infinite sum 1 + 1/2 + 1/4 + 1/8 + ...?", + "How many ways to arrange the letters in MISSISSIPPI?", + "What is the remainder when dividing 10^100 by 7?", + "Is sqrt(2) + sqrt(3) rational?", + "What is the number of partitions of 7?", + "How many edges in a complete bipartite graph K3,4?", + "What is the 50th term of the Fibonacci sequence?", + "Is 1 a prime number?", + "What is the sum of the first 100 positive integers?", + "How many surjections from a 4-set to a 3-set?", + "What is the value of the continued fraction [1;1,1,1,...]?", + "Is every integer either even or odd?", + "What is the area of a regular hexagon with side 1?", + "How many binary strings of length 10 have exactly 5 ones?", +] + +# ─── CODE: SIMPLE-classified prompts ─── +# These are short enough or lack code keywords to classify as SIMPLE + +# Genuinely trivial +CODE_SIMPLE_EASY = [ + "What does len([1,2,3]) return?", + "What is True and False in Python?", + "What does 'hello'[0] return?", + "What is type(42) in Python?", + "What does [1,2,3] + [4,5] give?", + "What is 10 // 3 in Python?", + "What does 'abc'.upper() return?", + "What is None == False in Python?", + "What does list(range(5)) produce?", + "What is 2 ** 10 in Python?", + "What does 'hello world'.split() return?", + "What is bool('') in Python?", + "What does [1,2,3][-1] return?", + "What is 5 % 3 in Python?", + "What does sorted([3,1,2]) return?", + "What is '' == False in Python?", + "What does {1,2,3} & {2,3,4} give?", + "What is type([]) in Python?", + "What does 'a' * 3 return?", + "What is 0.1 + 0.2 == 0.3 in Python?", + "What does len({}) return?", + "What is 3 > 2 > 1 in Python?", + "What does reversed([1,2,3]) give (as list)?", + "What is int('42') in Python?", + "What does 'hello'[-2:] return?", + "What is 1 == 1.0 in Python?", + "What does max(1, 2, 3) return?", + "What is [] == False in Python?", + "What does tuple([1,2,3]) return?", + "What is divmod(17, 5) in Python?", + "What does 'hello'.replace('l', 'r') return?", + "What is not not True in Python?", + "What does sum([1,2,3,4,5]) return?", + "What is chr(65) in Python?", + "What does [x**2 for x in range(4)] produce?", + "What is abs(-7) in Python?", + "What does 'abc' in 'abcdef' evaluate to?", + "What is round(2.5) in Python?", + "What does zip([1,2],[3,4]) give (as list)?", + "What is all([True, True, False]) in Python?", + "What does ''.join(['a','b','c']) return?", + "What is any([False, False, True]) in Python?", + "What does dict(a=1, b=2) produce?", + "What is isinstance(3, int) in Python?", + "What does [1,2,3].pop() return?", + "What is min('hello') in Python?", + "What does set([1,1,2,2,3]) give?", + "What is 'abc'.find('b') in Python?", + "What does enumerate(['a','b']) give (as list)?", + "What is pow(2, 3, 5) in Python?", + "What does 'Hello World'.title() return?", + "What is hash(42) == hash(42.0) in Python?", + "What does list('abc') return?", + "What is (1,2,3)[1:] in Python?", + "What does bin(10) return?", + "What is float('inf') > 10**100 in Python?", + "What does [None] * 3 produce?", + "What is 'ab' < 'ac' in Python?", + "What does hex(255) return?", + "What is [] is [] in Python?", + # 60 easy code prompts +] + +# Deceptively simple code: short, no/few code keywords, but tricky +CODE_SIMPLE_HARD = [ + "What does this print: a=[1,2,3]; b=a; b[0]=9; print(a)?", + "What is the output: x=5; print(x:=3, x)?", + "What does [i for i in range(10) if i%2 if i%3] produce?", + "What does (lambda f: f(f))(lambda f: 1) return?", + "What is the output: print(0.1 + 0.2 == 0.30000000000000004)?", + "What does this give: {True: 'a', 1: 'b', 1.0: 'c'}?", + "What is the output: x=[]; x.append(x); print(len(x))?", + "What does print(type(type)) output?", + "What is the result: (1,) + (1,) == (1, 1)?", + "What does print(1 << 20) give?", + "What is the output: a='hello'; a[0]='H'?", + "What does bool([False]) evaluate to?", + "What is the output: print({} == set())?", + "What does [*range(3), *range(3)] produce?", + "What is the output: x=256; y=256; print(x is y)?", + "What does print(round(0.5), round(1.5), round(2.5)) give?", + "What is the output: print('a' + 1)?", + "What does list(filter(None, [0,'',1,'a',False])) return?", + "What is the output: print(int(True) + int(False))?", + "What does print(f'{3.14:.10f}') show?", + "What is the result of: not not not True?", + "What does [1,2,3,4,5][::2] return?", + "What is the output: d={}; d[[1]]=1?", + "What does print(5 == 5.0 == True) give?", + "What is the output: print('abc'*0)?", + "What does next(iter([])) do?", + "What is the result: (True + True + True) * (True + True)?", + "What does {i:i for i in range(3)} produce?", + "What is the output: print(None is None)?", + "What does ' hello '.strip() == 'hello' evaluate to?", + "What is the result of max([], default=42)?", + "What does print(2**2**3) give?", + "What is the output: a=[1,2]; b=[1,2]; print(a==b, a is b)?", + "What does all([]) return in Python?", + "What is the output: print(ord('A'), ord('a'))?", + "What does frozenset([1,2,3]) == {1,2,3} evaluate to?", + "What is the result: print(-(-1))?", + "What does print(isinstance(True, int)) give?", + "What is the output: x=10; print(x if x > 5 else 'no')?", + "What does print(3 in [1,[3],2]) return?", +] + +# ─── CODE: MEDIUM-classified prompts ─── +# These have code keywords that push them into MEDIUM + +# Easy mediums: keywords trigger MEDIUM but the actual task is trivial +CODE_MEDIUM_EASY = [ + "Write a Python function that returns the string 'hello'.", + "Write a function to return the larger of two numbers.", + "Implement a function that checks if a number is even.", + "Write a Python function to convert Celsius to Fahrenheit.", + "Write a function that returns the last element of a list.", + "Implement a function to concatenate two strings with a space.", + "Write a function to check if a string is empty.", + "Write a Python function to compute absolute value without abs().", + "Implement a function that swaps two variables.", + "Write a function to check if a character is a vowel.", + "Write a function that doubles every element in a list.", + "Implement a function to count spaces in a string.", + "Write a function to return the first n elements of a list.", + "Write a Python function to compute the sum of a list.", + "Implement a function that reverses a string.", + "Write a function to find the minimum of three numbers.", + "Write a function to check if a year is a leap year.", + "Implement a function to remove all whitespace from a string.", + "Write a function to compute n factorial iteratively.", + "Write a Python function that returns True if a list is sorted.", + "Implement a function to count occurrences of a char in a string.", + "Write a function to check if two strings are equal ignoring case.", + "Write a function that returns a list of numbers from 1 to n.", + "Implement a function to compute the average of a list.", + "Write a function that flattens a list of lists.", + "Write a Python function to return the index of the max element.", + "Implement a function to repeat a string n times.", + "Write a function to remove duplicates from a list preserving order.", + "Write a function to check if all elements in a list are the same.", + "Implement a function to compute the dot product of two vectors.", + "Write a Python function to capitalize first letter of each word.", + "Write a function to count the number of words in a string.", + "Implement a function that chunks a list into sublists of size n.", + "Write a function to compute the running sum of a list.", + "Write a Python function to filter even numbers from a list.", + "Implement a function to check if a number is prime.", + "Write a function that returns a dict of character frequencies.", + "Write a function to compute Fibonacci(n) recursively.", + "Implement a function to find the longest string in a list.", + "Write a function to check if two lists have common elements.", + "Write a Python function to pad a string to a given length.", + "Implement a function that returns every other element of a list.", + "Write a function to find the most frequent element in a list.", + "Write a function to determine if a number is a power of two.", + "Implement a function to compute the Manhattan distance.", + "Write a Python function to check if a string is a palindrome.", + "Write a function to compute the Hamming distance between strings.", + "Implement a function to merge two sorted lists.", + "Write a function to find the second largest element in a list.", + "Write a Python function to check if parentheses are balanced.", + # 50 easy mediums +] + +# Hard mediums: correctly needs 9B capability +CODE_MEDIUM_HARD = [ + "Implement a linked list with insert, delete, and search methods in Python.", + "Write a function to perform an in-order traversal of a binary tree iteratively.", + "Implement a min-heap with insert and extract_min operations.", + "Write a function that solves the two-sum problem in O(n) time.", + "Implement a trie (prefix tree) with insert and search methods.", + "Write a function to find the longest palindromic substring in O(n^2).", + "Implement Dijkstra's shortest path algorithm for a weighted graph.", + "Write a function to serialize and deserialize a binary tree.", + "Implement an LRU cache with O(1) get and put operations.", + "Write a function to find all strongly connected components using Tarjan's algorithm.", + "Implement merge sort and verify it handles edge cases correctly.", + "Write a function to detect a cycle in a linked list and find the start of the cycle.", + "Implement a thread-safe bounded queue using locks.", + "Write a function to find the kth largest element in O(n) average time.", + "Implement a skip list with insert, search, and delete operations.", + "Write a function to compute the longest common subsequence of two strings.", + "Implement a red-black tree insertion with proper rotations and color fixes.", + "Write a function that evaluates a math expression string with +, -, *, / and parentheses.", + "Implement the A* pathfinding algorithm on a 2D grid.", + "Write a function to find all bridges in an undirected graph.", + "Implement a segment tree with range sum queries and point updates.", + "Write a function to solve the 0/1 knapsack problem using dynamic programming.", + "Implement topological sort using Kahn's algorithm.", + "Write a function to find the median of a data stream using two heaps.", + "Implement a Bloom filter with configurable false positive rate.", + "Write a function to solve the word break problem using DP.", + "Implement the KMP string matching algorithm.", + "Write a function to find the maximum flow using Ford-Fulkerson.", + "Implement a concurrent hash map with fine-grained locking.", + "Write a function to count inversions in an array using merge sort.", + "Implement a B-tree of order 3 with insert and search.", + "Write a function to solve N-Queens and return all valid configurations.", + "Implement a disjoint set (union-find) with path compression and union by rank.", + "Write a function to find the shortest path in a weighted DAG.", + "Implement a Fenwick tree (Binary Indexed Tree) for prefix sums.", + "Write a function to find all articulation points in an undirected graph.", + "Implement quicksort with three-way partitioning for arrays with duplicates.", + "Write a function to solve the edit distance problem using DP.", + "Implement a suffix array construction in O(n log n) time.", + "Write a function for matrix exponentiation to compute Fibonacci(n) in O(log n).", + "Implement a graph coloring algorithm using backtracking.", + "Write a function to find the maximum subarray sum with indices (Kadane's).", + "Implement a persistent stack data structure.", + "Write a function to solve the coin change problem (minimum coins) using DP.", + "Implement a Treap with insert, delete, and split operations.", + "Write a function to compute the convex hull of 2D points.", + "Implement the Bellman-Ford algorithm with negative cycle detection.", + "Write a function to find the longest path in a DAG.", + "Implement an interval tree with overlap queries.", + "Write a function to solve TSP for n<=15 using bitmask DP.", + # 50 hard code mediums +] + +# ─── CODE: Additional tricky SIMPLE prompts ─── +CODE_SIMPLE_TRICKY = [ + "What is the time complexity of list.append() in Python?", + "What does 'yield' do differently than 'return'?", + "What is the difference between == and 'is' in Python?", + "What does the GIL prevent in CPython?", + "What is the output of print(sys.getrecursionlimit()) typically?", + "What is a closure in Python? One sentence.", + "What happens if you modify a list while iterating over it?", + "What is the difference between deepcopy and shallow copy?", + "What does __slots__ do in a Python class?", + "What is the time complexity of 'in' for a Python set vs list?", + "What is the default hash of a user-defined class instance?", + "What does @staticmethod vs @classmethod do differently?", + "What is the MRO of a diamond inheritance in Python?", + "What happens when you divide by zero in Python?", + "What is the difference between a generator and a list comprehension?", + "What does the 'nonlocal' keyword do?", + "What is Python's garbage collection strategy?", + "What does __init__ vs __new__ do?", + "What is the result of sorting [None, 1, 'a'] in Python 3?", + "What is the default encoding in Python 3?", + "What does *args and **kwargs mean in a function signature?", + "What is a descriptor protocol in Python?", + "What does the 'with' statement guarantee?", + "What is the difference between @property and a regular method?", + "What happens if __hash__ is None for a class?", + "What is the difference between str and bytes in Python 3?", + "What does asyncio.gather() do vs asyncio.wait()?", + "What is the maximum recursion depth in Python?", + "What does the walrus operator := do?", + "What is the difference between a module and a package?", + "What is the result of float('nan') == float('nan')?", + "What does __repr__ vs __str__ do?", + "What is the difference between raise and raise from?", + "What does Python's 'finally' block guarantee?", + "What is the output of print(id(1000) == id(1000))?", + "What does a metaclass do in Python?", + "What is the difference between threading and multiprocessing?", + "What does sys.getsizeof([]) return approximately?", + "What is weak reference in Python used for?", + "What is the difference between ABC and Protocol?", + "What does collections.defaultdict(list) do differently than {}?", + "What is the output of bool(float('nan'))?", + "What does functools.lru_cache do?", + "What is the difference between Queue and deque?", + "What happens when you pickle a lambda?", + "What is the time complexity of dict.get() in Python?", + "What does itertools.chain do?", + "What is the result of [] + () in Python?", + "What does object.__eq__ default to?", + "What is the difference between copy() and = for dicts?", +] + +# ─── Additional filler to reach exactly 1000 ─── +# More SIMPLE-classified but tricky math prompts +MATH_SIMPLE_EXTRA = [ + "What is the 7th triangular number?", + "How many zeros in the product 1*2*3*...*20?", + "What is the digital root of 9999?", + "Is the sum of two irrationals always irrational?", + "What is the smallest number with exactly 6 divisors?", + "How many palindromic numbers between 1 and 1000?", + "What is 111111 / 7?", + "Is 2^10 + 1 prime?", + "What is the next prime after 97?", + "How many squares on a standard chessboard (all sizes)?", + "What is the harmonic mean of 2 and 6?", + "Is 0! = 1? Why or why not?", + "What is the 15th prime number?", + "How many trailing zeros in 50!?", + "What is the sum of the first 50 odd numbers?", + "Is pi transcendental?", + "What is the largest palindrome product of two 2-digit numbers?", + "How many distinct prime factors does 2310 have?", + "What is the golden ratio to 5 decimal places?", + "Is every multiple of 6 also a multiple of both 2 and 3?", + "What is the number of derangements of 5 elements?", + "How many integer solutions to |x| + |y| = 5?", + "What is the Euler totient of 12?", + "Is the set of rationals countable?", + "What is the sum of all digits of all numbers from 1 to 100?", + "How many paths in a 4x4 grid from top-left to bottom-right (only right/down)?", + "What is the probability of getting exactly 3 heads in 5 fair coin flips?", + "Is 10^10 + 1 divisible by 11?", + "What is the 12th Fibonacci number?", + "How many ways to seat 4 people at a round table?", +] + +# ─── Build PROMPT_METADATA ─── + +ALL_PROMPTS = [] +PROMPT_METADATA = [] + + +def _add_prompts(prompts, category, tier, expected_difficulty): + """Add prompts to the global lists with metadata.""" + for prompt in prompts: + PROMPT_METADATA.append({ + "index": len(ALL_PROMPTS), + "category": category, + "tier": tier, + "expected_difficulty": expected_difficulty, + "prompt": prompt, + }) + ALL_PROMPTS.append(prompt) + + +# Math prompts (500 total) +_add_prompts(MATH_SIMPLE_EASY, "math", 1, "trivial") # 70 - genuinely easy, SIMPLE tier +_add_prompts(MATH_SIMPLE_HARD, "math", 3, "hard_for_4b") # 30 - deceptively simple +_add_prompts(MATH_SIMPLE_TRICKY, "math", 2, "tricky") # 50 - tricky SIMPLE-classified +_add_prompts(MATH_SIMPLE_EXTRA, "math", 2, "tricky") # 30 - more tricky SIMPLE +_add_prompts(MATH_MEDIUM_EASY, "math", 1, "trivial") # 50 - easy MEDIUM-classified +_add_prompts(MATH_MEDIUM_HARD, "math", 4, "hard_for_4b") # 50 - hard MEDIUM-classified +# Subtotal: 280 math prompts... need 220 more + +# Extra math to reach 500 +_MATH_SIMPLE_FILL = [ + "What is 123 + 456?", + "What is 15 * 15?", + "What is the perimeter of a square with side 5?", + "How many seconds in a day?", + "What is 7! / 5!?", + "What is 1000 / 8?", + "How many milliliters in 1 liter?", + "What is the sum of angles in a triangle?", + "What is 99 + 1?", + "What is the area of a circle with radius 1? Leave in terms of pi.", + "Convert 5 km to meters.", + "What is 100/25?", + "How many diagonals does a hexagon have?", + "What is 15/35 simplified?", + "What is 7 + 8 * 2?", + "What is 6/9 simplified?", + "What is 24/36 simplified?", + "Convert 3/4 to a decimal.", + "What is 0.125 as a fraction?", + "What is 17 - (-3)?", + "Convert 72 degrees F to Celsius.", + "What is the square root of 49?", + "What is the largest single-digit prime?", + "What is 8^2 - 3^2?", + "How many vertices does a hexagon have?", + "What is 2^5 / 2^3?", + "What is 10^3 / 10?", + "What is 144 / 12?", + "What is 3 * 3 * 3 * 3?", + "How many mm in 1 cm?", + "What is 1/4 + 1/4?", + "What is the perimeter of a triangle with sides 3, 4, 5?", + "How many hours in a week?", + "What is 50 * 50?", + "What is the area of a square with side 7?", + "How many sides does an octagon have?", + "What is 1/2 + 1/3?", + "What is 16 * 16?", + "How many feet in a mile?", + "What is the product of 7 and 13?", + "What is 2^8?", + "How many degrees in a full rotation?", + "What is 1000 - 1?", + "What is the radius if diameter is 14?", + "What is 3/8 as a decimal?", + "How many seconds in 10 minutes?", + "What is 20% of 500?", + "What is the surface area of a cube with side 4?", + "How many weeks in a year?", + "What is 18 * 18?", + "What is the GCD of 24 and 36?", + "How many minutes in 3 hours?", + "What is 9^2 + 12^2?", + "What is the mean of {1, 2, 3, 4, 5}?", + "How many cm in 1 meter?", + "What is 7 * 7 * 7?", + "What is 1/5 as a percentage?", + "How many grams in a kilogram?", + "What is 15^2?", + "What is the volume of a sphere with radius 1? Leave in terms of pi.", + "How many days in February in a non-leap year?", + "What is 1024 / 2?", + "What is the largest 2-digit prime?", + "How many bits in a byte?", + "What is 6 * 7 * 8?", + "What is pi rounded to 4 decimal places?", + "How many ounces in a pound?", + "What is 2^16?", + "What is the mode of {1, 2, 2, 3, 3, 3}?", + "How many edges does a triangular prism have?", +] +_add_prompts(_MATH_SIMPLE_FILL[:70], "math", 1, "trivial") # 70 more easy filler + +_MATH_MEDIUM_FILL_EASY = [ + "Using the algorithm for computing averages, find the mean of [10, 20, 30, 40, 50].", + "Implement the formula for the area of a trapezoid with bases 3 and 7, height 4.", + "What is the performance of binary search compared to linear search for 1000 items?", + "Using the optimization technique, simplify the expression 2(x+3) - 2x.", + "Debug this: if a function returns None, what does bool(None) evaluate to?", + "Implement the formula for the perimeter of an ellipse with a=5, b=3 (approximate).", + "What is the throughput if a network sends 1000 packets in 2 seconds?", + "Using the algorithm for bubble sort, sort [5, 3, 8, 1].", + "Optimize: what is 4 * 25 * 17 using associativity?", + "Implement the algorithm to compute the Nth triangle number for N=10.", + "What is the latency of accessing L1 cache vs main memory (approximate ratio)?", + "Using the formula for combinations, compute C(6, 2).", + "Debug: what is the output of abs(-5) + abs(5)?", + "Implement the algorithm to check if 153 is a narcissistic number.", + "What is the memory required for a 1000x1000 matrix of 8-byte floats?", + "Using the algorithm for Euclidean GCD, compute GCD(56, 42).", + "Optimize: compute 2^10 without using exponentiation (bit shift).", + "Implement the formula for the volume of a cylinder with r=3, h=7.", + "What is the throughput of a CPU at 3 GHz executing 1 instruction per cycle?", + "Using the algorithm for insertion sort, sort [4, 2, 7, 1, 3].", + "Debug: what does 10 % 3 equal in Python vs in some other languages?", + "Implement the formula to convert 45 degrees to radians.", + "What is the performance cost of a cache miss?", + "Using the algorithm for computing powers, what is 2^15 via repeated squaring?", + "Optimize: is it faster to check divisibility by 2 using modulo or bitwise AND?", + "Implement the formula for the surface area of a cylinder with r=2, h=5.", + "What is the latency difference between SSD and HDD access in milliseconds?", + "Using the algorithm for selection sort, sort [9, 5, 2, 7, 3].", + "Debug: what does print(1/3) show in Python 3?", + "Implement the formula for converting binary 11001 to decimal.", + "What is the memory overhead of a Python dict compared to a list for 100 elements?", + "Using the algorithm for counting sort on [3, 1, 4, 1, 5, 9, 2, 6].", + "Optimize: is string concatenation with + or join faster for 1000 strings?", + "Implement the formula for the area of a regular pentagon with side 4.", + "What is the throughput limit of a single-threaded Python program?", + "Using the algorithm for radix sort, sort [170, 45, 75, 90, 802, 24, 2, 66].", + "Debug: what does 2 ** -1 evaluate to in Python?", + "Implement the formula for compound interest: $1000, 5%, 3 years, monthly.", + "What is the latency of a DNS lookup in milliseconds (typical)?", + "Using the algorithm for heap sort, sort [4, 10, 3, 5, 1].", + "Optimize: is multiplication or bit shifting faster for multiplying by 2?", + "Implement the formula for the distance between parallel lines y=2x+1 and y=2x+5.", + "What is the performance of Python list.sort() compared to sorted()?", + "Using the algorithm for merge sort, sort [38, 27, 43, 3, 9, 82, 10].", + "Debug: what is the difference between / and // in Python for negative numbers?", + "Implement the formula for converting hex FF to decimal.", + "What is the memory usage of a Python int for very large numbers?", + "Using the algorithm for computing factorials iteratively, compute 10!.", + "Optimize: which is more memory efficient, a tuple or a list in Python?", + "Implement the formula for the nth term of an arithmetic sequence: a1=3, d=5, n=10.", + "What is the throughput of gigabit ethernet in MB/s?", +] +_add_prompts(_MATH_MEDIUM_FILL_EASY[:50], "math", 1, "trivial") # 50 easy mediums +_add_prompts(_MATH_MEDIUM_FILL_EASY[50:], "math", 1, "trivial") # 1 more to round (use only 50) + +# Code prompts (500 total) +_add_prompts(CODE_SIMPLE_EASY, "code", 1, "trivial") # 60 - genuinely easy, SIMPLE tier +_add_prompts(CODE_SIMPLE_HARD, "code", 3, "hard_for_4b") # 40 - deceptively simple code +_add_prompts(CODE_SIMPLE_TRICKY, "code", 2, "tricky") # 50 - tricky SIMPLE-classified +_add_prompts(CODE_MEDIUM_EASY, "code", 1, "trivial") # 50 - easy MEDIUM-classified +_add_prompts(CODE_MEDIUM_HARD, "code", 4, "hard_for_4b") # 50 - hard MEDIUM-classified + +# Need 250 more code prompts (easy filler to balance) +_CODE_SIMPLE_FILL = [ + "What does str(123) return?", + "What is the type of 3.14 in Python?", + "What does [1,2,3].index(2) return?", + "What is 10 ** 0 in Python?", + "What does 'hello'.startswith('he') return?", + "What is bool(1) in Python?", + "What does [1,2,3].count(2) return?", + "What is type(None) in Python?", + "What does 'abcdef'[2:4] return?", + "What is 7 // 2 in Python?", + "What does len('hello') return?", + "What is 'a' + 'b' in Python?", + "What does [3,1,2].sort() return?", + "What is max(1, 5, 3) in Python?", + "What does 'hello'.endswith('lo') return?", + "What is bool(None) in Python?", + "What does [1,2] * 2 produce?", + "What is type(True) in Python?", + "What does '123'.isdigit() return?", + "What is 10 % 7 in Python?", + "What does 'Hello'.lower() return?", + "What is float(5) in Python?", + "What does [1,2,3,4][1:3] return?", + "What is str(True) in Python?", + "What does 'abc'.count('a') return?", + "What is int(3.9) in Python?", + "What does [1,2,3].insert(1, 'a') do to the list?", + "What is 2 + 3.0 in Python (type)?", + "What does 'hello world'.count('l') return?", + "What is len(set([1,1,2,2,3])) in Python?", + "What does 'abc'[::-1] return?", + "What is 5 > 3 in Python?", + "What does [1,2,3].extend([4,5]) do?", + "What is bool(0.0) in Python?", + "What does 'hello'.capitalize() return?", + "What is type({}) in Python?", + "What does list(reversed([1,2,3])) return?", + "What is 4 & 6 in Python?", + "What does 'hi there'.split(' ') return?", + "What is 4 | 3 in Python?", + "What does 'abc'.isalpha() return?", + "What is type(()) in Python?", + "What does (1,2) + (3,) return?", + "What is 8 >> 1 in Python?", + "What does 'test'.zfill(6) return?", + "What is 3 << 2 in Python?", + "What does ord('a') return?", + "What is hex(16) in Python?", + "What does chr(97) return?", + "What is oct(8) in Python?", + "What does 'abc'.replace('b', 'B') return?", + "What is 5 ^ 3 in Python?", + "What does 'hello'*2 return?", + "What is abs(-3.14) in Python?", + "What does [x for x in range(5) if x > 2] produce?", + "What is type(lambda x: x) in Python?", + "What does min([5, 2, 8, 1]) return?", + "What is 'abc' == 'ABC' in Python?", + "What does {1, 2} | {2, 3} give?", + "What is 1j * 1j in Python?", + "What does 'hello'.rjust(10) produce?", + "What is divmod(10, 3) in Python?", + "What does 'aabbcc'.count('bb') return?", + "What is round(3.14159, 2) in Python?", + "What does ' '.join(['a', 'b', 'c']) return?", + "What is bool(-1) in Python?", + "What does [i**2 for i in range(5)] produce?", + "What is 'hello'.index('l') in Python?", + "What does sum(range(10)) return?", + "What is type(range(5)) in Python?", +] +_add_prompts(_CODE_SIMPLE_FILL, "code", 1, "trivial") # 70 simple filler + +_CODE_MEDIUM_FILL_EASY = [ + "Write a Python function to return True if a number is positive.", + "Implement a function that takes a string and returns its length.", + "Write a function to add two numbers and return the result.", + "Implement a function to check if a list is empty.", + "Write a Python function to return the first element of a list.", + "Implement a function to multiply all elements in a list.", + "Write a function that returns True if a string contains a digit.", + "Implement a function to find the index of an element in a list.", + "Write a Python function to return the last character of a string.", + "Implement a function to count how many times 'a' appears in a string.", + "Write a function to return the maximum of a list without using max().", + "Implement a function to check if a number is between 1 and 100.", + "Write a Python function to create a list of zeros of length n.", + "Implement a function to swap the first and last elements of a list.", + "Write a function to return all elements greater than 5 from a list.", + "Implement a function to compute the remainder of a divided by b.", + "Write a Python function that returns the keys of a dictionary as a list.", + "Implement a function to check if a string starts and ends with the same char.", + "Write a function to return the sum of digits of a number.", + "Implement a function to triple every element in a list.", + "Write a Python function to return the middle element of a list.", + "Implement a function to check if a list has an even number of elements.", + "Write a function to return the unique elements of a list as a set.", + "Implement a function to concatenate a list of strings into one string.", + "Write a Python function that returns True if n is divisible by both 3 and 5.", + "Implement a function to return the smallest element of a list.", + "Write a function to check if two numbers have the same sign.", + "Implement a function that takes a dict and returns the number of keys.", + "Write a Python function to square each element in a list.", + "Implement a function to check if a string is all lowercase.", + "Write a function to return the values of a dictionary as a list.", + "Implement a function to insert an element at the beginning of a list.", + "Write a Python function to return the second element of a tuple.", + "Implement a function to check if a list contains only positive numbers.", + "Write a function to return a list without its first element.", + "Implement a function to convert a list of ints to a list of strings.", + "Write a Python function to return the intersection of two sets.", + "Implement a function to check if a character is uppercase.", + "Write a function to append an element to a list and return it.", + "Implement a function that returns the type of its argument as a string.", + "Write a Python function to return a sorted copy of a list.", + "Implement a function to check if a number is a multiple of 7.", + "Write a function to return the number of vowels in a string.", + "Implement a function to create a dictionary from two lists (keys and values).", + "Write a Python function to return True if all elements in a list are positive.", + "Implement a function to remove the first occurrence of a value from a list.", + "Write a function to return the absolute difference of two numbers.", + "Implement a function to check if a string is numeric.", + "Write a Python function to return a list in reverse order without reverse().", + "Implement a function to count the number of uppercase letters in a string.", + "Write a function to return the product of the first n natural numbers.", + "Implement a function to check if two strings are the same length.", + "Write a Python function to return the last n elements of a list.", + "Implement a function to find the index of the minimum element.", + "Write a function to check if a number is a perfect cube.", + "Implement a function that returns True if a list is a subset of another.", + "Write a Python function to merge two dictionaries.", + "Implement a function to return the factorial of n using recursion.", + "Write a function to check if a string reads the same forwards and backwards.", + "Implement a function to return the n largest elements of a list.", + "Write a Python function to compute the power of a number using a loop.", + "Implement a function to flatten a list that contains sublists one level deep.", + "Write a function to return the number of even numbers in a list.", + "Implement a function to check if a year is in the 21st century.", + "Write a Python function to return every third element from a list.", + "Implement a function to compute the distance between two points (x1,y1) and (x2,y2).", + "Write a function to check if a list contains duplicates.", + "Implement a function to return the common elements of three lists.", + "Write a Python function to generate a list of the first n even numbers.", + "Implement a function to check if a number is negative.", + "Write a function that returns the longer of two strings.", + "Implement a function to compute the average of the top 3 scores in a list.", + "Write a Python function to return a string with all vowels removed.", + "Implement a function to check if a list is in descending order.", + "Write a function to return the symmetric difference of two sets.", + "Implement a function to capitalize only the first character of a string.", + "Write a Python function to return True if a number rounds to 10.", + "Implement a function to split a list into two halves.", + "Write a function to return the second-to-last element of a list.", + "Implement a function that returns 'fizz' if n%3==0, 'buzz' if n%5==0, else n.", +] +_add_prompts(_CODE_MEDIUM_FILL_EASY, "code", 1, "trivial") # 80 easy medium filler + +# ─── More prompts to reach 1000 ─── + +# More deceptively simple math (SIMPLE-classified, hard for 4B) +_MATH_SIMPLE_HARD_2 = [ + "What is the sum of all divisors of 28?", + "How many distinct ways to make 50 cents using US coins?", + "What is the maximum number of regions formed by 5 lines in a plane?", + "Is the square root of 2 plus the square root of 3 greater than 3?", + "What is the largest prime gap below 100?", + "How many integers from 1 to 100 are not divisible by 2, 3, or 5?", + "What is the remainder when 3^100 is divided by 4?", + "Is 341 a pseudoprime base 2?", + "What is the sum of the infinite series 1 - 1/2 + 1/3 - 1/4 + ...?", + "How many permutations of {1,2,3,4,5} have no fixed points?", + "What is the probability of getting at least one 6 in 4 dice rolls?", + "Is the sum 1/2 + 1/3 + 1/5 + 1/7 + 1/11 greater than 1?", + "What is the smallest positive integer with exactly 12 divisors?", + "How many lattice points are inside the circle x^2 + y^2 < 10?", + "What is (-1)^(1/2) in the complex numbers?", + "Is 2^89 - 1 prime?", + "What is the expected number of coin flips to get two heads in a row?", + "How many connected graphs on 4 labeled vertices exist?", + "What is the chromatic number of the complete graph K4?", + "Is e^(i*pi) + 1 = 0? Why?", + "What is the number of distinct binary trees with 4 nodes?", + "How many ways to tile a 2x6 grid with 1x2 dominoes?", + "What is the exact value of cos(pi/5)?", + "Is 2^11 - 1 = 2047 prime? If not, what are its factors?", + "What is the Collatz sequence starting from 27? How many steps to reach 1?", + "How many integers between 1 and 1000 are perfect powers?", + "What is the value of sum k=1 to 100 of (-1)^k * k?", + "Is the number 111111111 (nine 1s) prime?", + "What is the smallest number expressible as sum of two cubes in two ways?", + "How many graphs on 3 labeled vertices exist?", +] +_add_prompts(_MATH_SIMPLE_HARD_2, "math", 3, "hard_for_4b") # 30 + +# More hard MEDIUM math +_MATH_MEDIUM_HARD_2 = [ + "Using the algorithm for computing determinants, find det([[1,2,3],[4,5,6],[7,8,10]]).", + "Implement the algorithm to find the shortest path in a weighted graph: A-B:3, B-C:1, A-C:10, B-D:2, C-D:7.", + "Using the algorithm for modular exponentiation, compute 7^256 mod 13.", + "Implement the algorithm for computing the convex hull of points: (0,0),(1,4),(3,1),(4,3),(2,2).", + "Using dynamic programming, find the number of ways to partition 15 into positive integers.", + "Implement the algorithm for Strassen's matrix multiplication on 2x2 matrices.", + "Using the algorithm for Prim's MST on edges: A-B:7, A-D:5, B-C:8, B-D:9, C-D:6, C-E:3, D-E:15.", + "Implement the algorithm to solve the fractional knapsack: items [(w=10,v=60),(w=20,v=100),(w=30,v=120)], capacity=50.", + "Using the algorithm for BFS, find the shortest path from A to F in: A-B, A-C, B-D, C-D, D-E, E-F, C-F.", + "Implement the algorithm for Kruskal's MST on: A-B:4, A-H:8, B-H:11, B-C:8, C-D:7, C-F:4, D-E:9, D-F:14.", + "Using the algorithm for computing eigenvalues via the characteristic polynomial of [[3,1],[1,3]].", + "Implement the algorithm to check if a graph is bipartite: edges 1-2, 2-3, 3-4, 4-1, 1-3.", + "Using the algorithm for Floyd's cycle detection, find the cycle in: 1->2->3->4->5->3.", + "Implement the algorithm for maximum matching in: A-1, A-2, B-1, B-3, C-2, C-3.", + "Using the algorithm for computing the rank of matrix [[1,2,3],[4,5,6],[5,7,9]].", + "Implement the algorithm for computing all permutations of [1,2,3,4] and count them.", + "Using the algorithm for interval scheduling, find max non-overlapping intervals from: [1,3],[2,5],[4,7],[6,9],[8,10].", + "Implement the algorithm for computing the power set of {a,b,c,d}. How many elements?", + "Using the algorithm for BFS on a grid, find shortest path from (0,0) to (3,3) avoiding (1,1),(2,2).", + "Implement the algorithm for the stable marriage problem with preferences: M1:[W1,W2], M2:[W2,W1]; W1:[M2,M1], W2:[M1,M2].", + "Using the algorithm for computing binomial coefficients, fill Pascal's triangle row 7.", + "Implement the algorithm for finding the median of [3,1,4,1,5,9,2,6,5,3,5] in O(n) average.", + "Using the algorithm for the Dutch national flag problem, partition [2,0,1,2,0,1,0,2,1] into [0s,1s,2s].", + "Implement the algorithm for computing the Levenshtein distance between 'intention' and 'execution'.", + "Using the algorithm for reservoir sampling, explain how to sample 5 items uniformly from a stream of unknown length.", + "Implement the algorithm for computing the nth row of Pascal's triangle for n=10.", + "Using the algorithm for topological sort on: CS101->CS201, CS101->CS202, CS201->CS301, CS202->CS301, CS202->CS303.", + "Implement the algorithm for Kadane's maximum subarray on [-2,1,-3,4,-1,2,1,-5,4]. Return max sum and indices.", + "Using the algorithm for counting sort, sort [4,2,2,8,3,3,1,7,4,2] and state the time complexity.", + "Implement the algorithm for the activity selection problem: activities [(1,4),(3,5),(0,6),(5,7),(3,9),(5,9),(6,10),(8,11),(8,12),(2,14)].", + "Using the algorithm for computing Catalan numbers, find C(7) and explain what it counts.", + "Implement the algorithm for detecting negative cycles using Bellman-Ford on: A->B:1, B->C:-3, C->A:1.", + "Using the algorithm for LCS, find the longest common subsequence of 'AGGTAB' and 'GXTXAYB'.", + "Implement the algorithm for 0/1 knapsack: items [(w=1,v=1),(w=3,v=4),(w=4,v=5),(w=5,v=7)], capacity=7.", + "Using the algorithm for graph coloring, find the chromatic number of C5 (5-cycle).", + "Implement the algorithm for computing the number of inversions in [8,4,2,1].", + "Using the algorithm for binary indexed tree, compute prefix sums of [3,2,4,5,1,6,2,8] up to each index.", + "Implement the algorithm for the coin row problem: coins [5,1,2,10,6,2]. Maximum value without adjacent?", + "Using the algorithm for string matching (KMP), find all occurrences of 'aba' in 'abababababa'.", + "Implement the algorithm for computing the number of paths in a 5x5 grid with obstacles at (2,2) and (3,3).", +] +_add_prompts(_MATH_MEDIUM_HARD_2, "math", 4, "hard_for_4b") # 40 + +# More hard code MEDIUM prompts +_CODE_MEDIUM_HARD_2 = [ + "Write a function to find the longest substring without repeating characters.", + "Implement a function to check if a binary tree is balanced.", + "Write a function to find all paths from root to leaves in a binary tree.", + "Implement a function to rotate a matrix 90 degrees clockwise in-place.", + "Write a function to find the intersection point of two linked lists.", + "Implement a function to convert a sorted array to a balanced BST.", + "Write a function to find the lowest common ancestor of two nodes in a BST.", + "Implement a function to clone a graph (deep copy of a graph with cycles).", + "Write a function to implement a stack that supports getMin() in O(1).", + "Implement a function to find the number of islands in a 2D grid.", + "Write a function to check if a graph is a valid tree.", + "Implement a function to find the kth smallest element in a BST.", + "Write a function to design a data structure supporting insert, delete, and getRandom in O(1).", + "Implement a function to find all valid combinations of n pairs of parentheses.", + "Write a function to merge k sorted linked lists into one sorted list.", + "Implement a function to find the diameter of a binary tree.", + "Write a function to determine if a Sudoku board is valid.", + "Implement a function to find the trap water problem solution given elevation map.", + "Write a function to compute the power(x, n) handling negative exponents.", + "Implement a function to find the next permutation of a number array.", + "Write a function to solve the jump game (can you reach the last index?).", + "Implement a function to find the minimum window containing all chars of pattern.", + "Write a function to compute the maximum profit from at most 2 stock transactions.", + "Implement a function to build a trie and find all words matching a pattern with '.'.", + "Write a function to find the largest rectangle containing only 1s in a binary matrix.", + "Implement a function to solve the word ladder problem (shortest transformation).", + "Write a function to compute the skyline of a set of buildings.", + "Implement a function to serialize and deserialize an N-ary tree.", + "Write a function to find the critical connections (bridges) in a network graph.", + "Implement a function to solve the regex matching problem with '.' and '*'.", +] +_add_prompts(_CODE_MEDIUM_HARD_2, "code", 4, "hard_for_4b") # 30 + +# More tricky code SIMPLE prompts +_CODE_SIMPLE_TRICKY_2 = [ + "What is the output: print([1,2,3] == [1,2,3], [1,2,3] is [1,2,3])?", + "What does print(sum(range(101))) output?", + "What is the result of 'hello'[1:4:2]?", + "What does print(2**2**2**2) give?", + "What is the output: x=[1]; x*=3; x[0]=9; print(x)?", + "What does print(len(set('mississippi'))) return?", + "What is the output of [1,2,3][3:]?", + "What does print(0 or '' or [] or 'hello' or 42) give?", + "What is the result of True + True + True?", + "What does print(sorted([3,1,2], reverse=True)) output?", + "What is the output: a=(1,); b=(1,); print(a==b, a is b)?", + "What does print(list(map(str, [1,2,3]))) give?", + "What is the result of 'abc'[::2]?", + "What does print(bool(float('inf'))) return?", + "What is the output: d={1:'a',2:'b'}; print(d.get(3,'default'))?", + "What does print(type(1/1)) give in Python 3?", + "What is the result of [*'hello']?", + "What does print(3 * 'ab' == 'ababab') return?", + "What is the output of print(None or 0 or '' or False or 'yes')?", + "What does print(list(zip('abc', [1,2]))) give?", + "What is the result: x={}; x[0]=x; print(type(x[0]))?", + "What does print(2 in [1, [2], 3]) return?", + "What is the output: a=[1,2]; b=a[:]; b.append(3); print(a)?", + "What does print('' in 'hello') return?", + "What is the result of max('hello')?", + "What does print(1_000_000) output in Python?", + "What is the output: print(0 == False, 0 is False)?", + "What does print(len(range(10,0,-2))) return?", + "What is the result of (1,2,3)*2?", + "What does print(all('')) return?", + "What is the output: x='hello'; print(x[-1:-4:-1])?", + "What does print(type(...)) return in Python?", + "What is the result of {1:2, 3:4}.values()?", + "What does print(10 == 10.0 == 10+0j) return?", + "What is the output of print(list(enumerate('ab', start=1)))?", + "What does print(dict(zip('abc', range(3)))) give?", + "What is the result: x=[0]*3; x[0]=[1]; print(x)?", + "What does print(min(None, 1)) do?", + "What is the output: print(f'{255:08b}')?", + "What does print(complex(1,2) + complex(3,4)) return?", + "What is the result of 'hello'.partition('l')?", + "What does print(set() == frozenset()) return?", + "What is the output: a=[1]; b=[1]; a.extend(b); print(a, b)?", + "What does print((lambda: 42)()) give?", + "What is the result of bytes(3)?", + "What does print({True: 1, 1: 2, 1.0: 3}) give?", + "What is the output: print([] < [1])?", + "What does print('abc'.maketrans('abc', 'xyz')) return?", + "What is the result of print(not not not False)?", +] +_add_prompts(_CODE_SIMPLE_TRICKY_2, "code", 2, "tricky") # 49 + +# More easy code SIMPLE filler +_CODE_SIMPLE_FILL_2 = [ + "What is str(None) in Python?", + "What does [1,2,3,4,5][:3] return?", + "What is 9 ** 0.5 in Python?", + "What does 'python'.upper() return?", + "What is type(3+4j) in Python?", + "What does [0] * 5 produce?", + "What is 'hello'[4] in Python?", + "What does bool([0]) return?", + "What is int('0b101', 2) in Python?", + "What does 'hello world'.title() return?", + "What is 15 & 9 in Python?", + "What does list('hello') return?", + "What is 15 | 9 in Python?", + "What does tuple('abc') return?", + "What is 100 // 7 in Python?", + "What does 'Hi'.swapcase() return?", + "What is type(b'hello') in Python?", + "What does [1,2,3].pop(0) return?", + "What is float('1e3') in Python?", + "What does 'abcabc'.rfind('c') return?", +] +_add_prompts(_CODE_SIMPLE_FILL_2, "code", 1, "trivial") # 20 + +# 30 more easy SIMPLE math to reach 1000 +_MATH_SIMPLE_FINAL = [ + "What is 25 + 75?", + "What is 6 * 6?", + "How many legs does a spider have?", + "What is 1/2 of 100?", + "What is 8 + 8 + 8?", + "How many sides does a triangle have?", + "What is 10 * 10 * 10?", + "What is 500 / 5?", + "How many minutes in an hour?", + "What is 33 + 67?", + "What is 2 * 2 * 2 * 2?", + "How many months in a year?", + "What is 1000 / 10?", + "What is 9 + 9 + 9?", + "How many days in a week?", + "What is 7 * 7?", + "What is 100 / 4?", + "How many zeros in a billion?", + "What is 12 + 12?", + "What is 3 * 5 * 7?", + "How many letters in the English alphabet?", + "What is 60 / 12?", + "What is 11 + 22 + 33?", + "How many planets in our solar system?", + "What is 8 * 9?", + "What is 200 / 4?", + "How many colors in a rainbow?", + "What is 15 + 25?", + "What is 4 * 4 * 4?", + "How many continents are there?", +] +_add_prompts(_MATH_SIMPLE_FINAL, "math", 1, "trivial") # 30 + +# Final count check and trim +if len(PROMPT_METADATA) > 1000: + PROMPT_METADATA = PROMPT_METADATA[:1000] + ALL_PROMPTS = ALL_PROMPTS[:1000] + diff --git a/scripts/test_custom_router_hybrid_v1.py b/scripts/test_custom_router_hybrid_v1.py new file mode 100644 index 00000000000..e20c977e310 --- /dev/null +++ b/scripts/test_custom_router_hybrid_v1.py @@ -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()) diff --git a/tests/test_litellm/router_strategy/test_hybrid_router.py b/tests/test_litellm/router_strategy/test_hybrid_router.py new file mode 100644 index 00000000000..0eee17bcfb1 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_hybrid_router.py @@ -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",)