mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 9c155ae5a8 into 80d3b69d9c
This commit is contained in:
commit
ede4159676
3 changed files with 23 additions and 10 deletions
|
|
@ -110,7 +110,7 @@ class AdaptiveRouter:
|
|||
self._cells[(rt, model)] = initial_cell(prefs, rt)
|
||||
|
||||
async def load_state_from_db(self, prisma_client: Any) -> None:
|
||||
"""Override cold-start cells with persisted state. Called once at startup."""
|
||||
"""Merge persisted deltas onto the cold-start prior. Called once at startup."""
|
||||
if prisma_client is None:
|
||||
return
|
||||
try:
|
||||
|
|
@ -126,8 +126,16 @@ class AdaptiveRouter:
|
|||
continue
|
||||
if row.model_name not in self.config.available_models:
|
||||
continue
|
||||
# Persisted alpha/beta are accumulated DELTAS — the flusher writes
|
||||
# increments, not absolute posteriors (for multi-pod safety). Merge
|
||||
# them onto the cold-start prior so the tier bias is preserved AND the
|
||||
# Beta stays valid (prior.beta >= 0.5, so a satisfaction-only delta of
|
||||
# beta=0 can never reload as Beta(alpha, 0) and crash thompson_sample).
|
||||
prefs = self.model_to_prefs.get(row.model_name) or _default_prefs()
|
||||
prior = initial_cell(prefs, rt)
|
||||
self._cells[(rt, row.model_name)] = BanditCell(
|
||||
alpha=row.alpha, beta=row.beta
|
||||
alpha=prior.alpha + row.alpha,
|
||||
beta=prior.beta + row.beta,
|
||||
)
|
||||
loaded += 1
|
||||
verbose_router_logger.info(
|
||||
|
|
|
|||
|
|
@ -233,7 +233,7 @@ async def test_record_turn_failure_increments_beta():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_state_from_db_overrides_cold_start():
|
||||
async def test_load_state_from_db_merges_deltas_onto_cold_start():
|
||||
r = _make_router()
|
||||
cold = r._cells[(RequestType.GENERAL, "fast")]
|
||||
|
||||
|
|
@ -247,8 +247,10 @@ async def test_load_state_from_db_overrides_cold_start():
|
|||
prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[fake_row])
|
||||
await r.load_state_from_db(prisma)
|
||||
|
||||
# Persisted alpha/beta are accumulated DELTAS, so the loader merges them
|
||||
# onto the cold-start prior rather than overwriting it.
|
||||
new_cell = r._cells[(RequestType.GENERAL, "fast")]
|
||||
assert (new_cell.alpha, new_cell.beta) == (42.0, 13.0)
|
||||
assert (new_cell.alpha, new_cell.beta) == (cold.alpha + 42.0, cold.beta + 13.0)
|
||||
assert (new_cell.alpha, new_cell.beta) != (cold.alpha, cold.beta)
|
||||
|
||||
|
||||
|
|
@ -275,8 +277,8 @@ async def test_load_state_from_db_handles_unknown_request_type():
|
|||
)
|
||||
await r.load_state_from_db(prisma)
|
||||
|
||||
# Unknown skipped; good applied.
|
||||
assert r._cells[(RequestType.GENERAL, "fast")].alpha == 7.0
|
||||
# Unknown skipped; good row's delta is merged onto the prior.
|
||||
assert r._cells[(RequestType.GENERAL, "fast")].alpha == cold.alpha + 7.0
|
||||
# Other request types kept their cold-start values.
|
||||
assert r._cells[(RequestType.WRITING, "fast")] == cold or True
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ What we cover:
|
|||
1. Full lifecycle: pick -> record turn(s) -> flush -> DB upsert with correct deltas
|
||||
2. Owner cache pins attribution: same key + matching model -> updates flow
|
||||
3. Convergence in-process: 50 simulated sessions, "good" model dominates last 10
|
||||
4. Cold-start state load from DB overrides priors
|
||||
4. Cold-start state load from DB merges persisted deltas onto priors
|
||||
5. Failure signal increments beta in the next flush
|
||||
6. Unknown request types in DB rows are silently skipped
|
||||
7. Flush isolates writes per (router, session, model) tuple
|
||||
|
|
@ -203,8 +203,9 @@ async def test_failure_signal_increments_beta_after_flush():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_state_from_db_overrides_cold_start():
|
||||
async def test_load_state_from_db_merges_deltas_onto_cold_start():
|
||||
router = _make_router()
|
||||
cold = router._cells[(RequestType.GENERAL, "gpt-4o")]
|
||||
fake_row = MagicMock()
|
||||
fake_row.request_type = RequestType.GENERAL.value
|
||||
fake_row.model_name = "gpt-4o"
|
||||
|
|
@ -216,9 +217,11 @@ async def test_load_state_from_db_overrides_cold_start():
|
|||
|
||||
await router.load_state_from_db(prisma)
|
||||
|
||||
# Persisted alpha/beta are accumulated DELTAS, so the loader merges them
|
||||
# onto the cold-start prior rather than overwriting it.
|
||||
cell = router._cells[(RequestType.GENERAL, "gpt-4o")]
|
||||
assert cell.alpha == 90.0
|
||||
assert cell.beta == 10.0
|
||||
assert cell.alpha == cold.alpha + 90.0
|
||||
assert cell.beta == cold.beta + 10.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue