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