sora 4eb853c8fc Gate ramp requires completions: zero throughput = hold, not +1
Ramping on demand alone was dangerous with slow endpoints: if the 2
in-flight longbench_v2 requests hang, 'no failures + waiters queued'
kept adding +1 every 5s all the way to 96 -- piling prefills onto a
server that had not answered anything. Both ramp paths (demand-driven
and /metrics) now require at least one SUCCESSFUL completion in the
probe interval; hangs hold the gate until read-timeouts fire and the
x0.7 backoff takes over.

Unit-verified: hang 4 probe intervals with 20 waiters -> limit stays 2;
one success -> +1; one failure -> x0.7.

Co-Authored-By: Claude <noreply@anthropic.com>
2026-09-15 02:33:21 +00:00

332 lines
14 KiB
Python

"""Multi-endpoint load balancing: one logical model, N local ports.
from evalharness.model.pool import PooledAdapter
from evalharness.model.adapter import resolve_adapter
base = resolve_adapter('openai/http://127.0.0.1:8123/v1?Qwen3-8B')
pool = PooledAdapter([resolve_adapter(f'openai/http://127.0.0.1:{p}/v1?Qwen3-8B')
for p in range(8123, 8131)])
out = await pool.generate(...) # round-robin over instances
Traffic management: round-robin keeps per-endpoint traffic even; backends
that fail repeatedly enter a cool-down window and are skipped until it
expires, so one sick endpoint cannot absorb its share of the load.
"""
import asyncio
import contextlib
import itertools
import time
from typing import Dict, List, Optional
from ..data.sample import ChatMessage
from .adapter import ModelAdapter
from .output import ModelOutput, Usage
class PooledAdapter(ModelAdapter):
"""Round-robin over N equivalent backend instances with health cooling."""
name = 'pool'
COOLDOWN_S = 60.0 # a backend that failed EVERY attempt rests this long
COOLDOWN_AFTER = 2 # consecutive full-pass failures before cooling
def __init__(self, adapters: List[ModelAdapter]):
if not adapters:
raise ValueError('PooledAdapter needs at least one backend')
super().__init__(model=adapters[0].model, api_base=adapters[0].api_base)
self.adapters = adapters
self._cycle = itertools.cycle(range(len(adapters)))
self.usage = Usage()
# request outcome counters (success rate accounting)
self.stats = {'requests': 0, 'ok': 0, 'failed': 0, 'retried': 0}
# per-backend health: consecutive_failures, cooling_until, per-endpoint counts
self._health = [{ 'fails': 0, 'until': 0.0, 'ok': 0, 'req': 0}
for _ in adapters]
# adaptive per-endpoint concurrency gates (AIMD over /metrics signals)
self._gates = [AdaptiveGate(a) for a in adapters]
def _next(self) -> ModelAdapter:
"""Round-robin, skipping endpoints inside their cool-down window."""
n = len(self.adapters)
now = time.time()
for _ in range(n):
i = next(self._cycle)
h = self._health[i]
if h['until'] <= now or all(x['until'] <= now for x in self._health):
self._health[i]['req'] += 1
return self.adapters[i]
# everything cooling: take the next anyway (better to try than stall)
i = next(self._cycle)
self._health[i]['req'] += 1
return self.adapters[i]
def _mark(self, adapter: ModelAdapter, ok: bool) -> None:
try:
i = self.adapters.index(adapter)
except ValueError:
return
h = self._health[i]
if ok:
h['fails'] = 0
h['until'] = 0.0
h['ok'] += 1
else:
h['fails'] += 1
if h['fails'] >= self.COOLDOWN_AFTER:
h['until'] = time.time() + self.COOLDOWN_S
h['fails'] = 0
def request_stats(self) -> Dict[str, float]:
"""Success-rate + per-endpoint traffic view (load-balance audit)."""
n = self.stats['requests']
out = {
'requests': n,
'success_rate': self.stats['ok'] / n if n else 0.0,
'retry_rate': self.stats['retried'] / n if n else 0.0,
'failure_rate': self.stats['failed'] / n if n else 0.0,
}
for i, (a, h) in enumerate(zip(self.adapters, self._health)):
tag = a.api_base.rsplit('//', 1)[-1].replace('/', '_')
out[f'ep{i}_{tag}_reqs'] = h['req']
out[f'ep{i}_{tag}_ok'] = h['ok']
if i < len(self._gates):
for k, v in self._gates[i].report().items():
out[f'ep{i}_{tag}_{k}'] = v
return out
async def generate(self, messages: List[ChatMessage],
tools: Optional[list] = None, **kw) -> ModelOutput:
last_exc = None
self.stats['requests'] += 1
for _ in range(len(self.adapters)): # try each instance once
adapter = self._next()
try:
i = self.adapters.index(adapter)
await self._gates[i].acquire()
ok = False
try:
out = await adapter.generate(messages, tools=tools, **kw)
ok = True
finally:
# ONE release per acquire: the old code released True in
# this finally AND False in the except handler, double-
# decrementing _inflight (gate over-admits) and applying
# the x0.7 backoff twice per failure
self._gates[i].release(ok)
self.usage = self.usage + out.usage
self.stats['ok'] += 1
self._mark(adapter, True)
if out.usage.retries:
self.stats['retried'] += 1
return out
except Exception as e: # dead/overloaded instance -> next
last_exc = e
self._mark(adapter, False)
# 4xx (e.g. 400 overloaded) still worth trying ANOTHER instance:
# one backend's state can differ from the rest
continue
self.stats['failed'] += 1
raise last_exc
async def close(self) -> None:
for a in self.adapters:
await a.close()
for g in self._gates:
g.stop()
class AdaptiveGate:
"""Per-endpoint adaptive concurrency limiter (AIMD + server signals).
Goal: keep the backend SATURATED (high XPU util / throughput) without
pushing it over the cliff (500s / child crashes). Signals:
- server /metrics: num_queue_reqs > 0 means WE are pushing too hard
for the current mix; idle (no queue, low running) means room to grow
- request failures: multiplicative decrease (survive first)
Control law (classic AIMD):
+1 concurrency per probe interval when the endpoint looks underfed
-1 when the server reports a queue (gentle)
x0.7 on any failed request (fast backoff), floor at LO
Purely additive to PooledAdapter: one gate per backend, no caller change.
"""
LO = 2 # never go below: progress beats perfection
HI = 96 # sane ceiling for one endpoint
PROBE_S = 5.0 # metrics probe interval
INITIAL = 8.0 # class-level start point (--auto-concurrency rebinds it)
def __init__(self, adapter: ModelAdapter):
self.adapter = adapter
self.limit = self.INITIAL # float for smooth x0.7; compare with int()
self._inflight = 0
self._cond: Optional[asyncio.Condition] = None
self._task: Optional[asyncio.Task] = None
self._stopped = False
self._loop = None # loop the cond/probe-task are bound to
self._contended = 0 # acquire() waits this probe interval (demand)
self._interval_fails = 0 # failed releases this probe interval
self._interval_ok = 0 # SUCCESSFUL releases this probe interval:
# zero completions = zero throughput, and a
# gate that ramps on demand alone would pile
# 96 concurrent prefills onto a server whose
# first 2 requests have not even answered
self.stats = {'probe': 0, 'ramp': 0, 'hold_queue': 0, 'backoff_fail': 0,
'backoff_queue': 0, 'ramp_demand': 0}
def _push_limit(self) -> None:
"""Surface the current limit to the progress bar ('gate N')."""
rep = (self.adapter.extra or {}).get('progress_reporter')
fn = getattr(rep, 'set_gate', None)
if fn is not None:
try:
fn(max(1, int(self.limit)))
except Exception:
pass
# ---- gate semantics ----
async def acquire(self) -> None:
loop = asyncio.get_running_loop()
if self._cond is None or self._loop is not loop or self._stopped:
# lazy init OR LOOP CHANGE: pools are cached across benchmarks,
# and the CLI runs asyncio.run() per bench (per repeat!) -- a new
# run means a new event loop while this gate survives. A Condition
# is loop-bound: reusing the old one raises "bound to a different
# event loop" under contention, and the old probe task died with
# the closed loop. Rebuild both; inflight resets to 0 (nothing is
# in flight on a fresh loop by construction).
self._loop = loop
self._cond = asyncio.Condition()
self._inflight = 0
self._stopped = False
self._push_limit()
if self._task is not None:
self._task.cancel() # dead task from the closed loop; no-op
self._task = loop.create_task(self._probe_loop())
while self._inflight >= max(1, int(self.limit)):
self._contended += 1 # demand exceeded the cap: potential ramp fuel
await self._cond.acquire()
try:
await self._cond.wait()
finally:
self._cond.release()
self._inflight += 1
def release(self, ok: bool) -> None:
self._inflight = max(0, self._inflight - 1)
if ok:
self._interval_ok += 1
else: # multiplicative decrease -- survival first
self._interval_fails += 1
before = self.limit
self.limit = max(self.LO, self.limit * 0.7)
if before != self.limit:
self.stats['backoff_fail'] += 1
self._push_limit()
self._wake()
def _wake(self) -> None:
if self._cond is not None:
# fire-and-forget notify (loop may not be ours -- best effort)
try:
fut = asyncio.ensure_future(self._notify_all())
fut.add_done_callback(lambda f: None)
except RuntimeError:
pass
async def _notify_all(self) -> None:
async with self._cond:
self._cond.notify_all()
# ---- server-signal probe ----
def _no_signal_ramp(self) -> None:
"""No server signals available (no /metrics, 404/HTTPError, gateway
stripped it, non-sglang backend): fall back to demand-driven AIMD --
ramp while the cap is the binding constraint (callers had to WAIT on
acquire) and the interval was failure-free. Failures still cut x0.7
per release, so a drowning backend shrinks the gate immediately."""
if self._interval_fails == 0 and self._interval_ok > 0 \
and self._contended > 0 and int(self.limit) < self.HI:
self.limit = min(self.HI, self.limit + 1)
self.stats['ramp_demand'] += 1
self._push_limit()
self._interval_fails = 0
self._interval_ok = 0
self._contended = 0
self._wake()
async def _probe_once(self) -> None:
import urllib.request
url = f'{self.adapter.api_base.rstrip("/")}/metrics'
def _fetch():
with urllib.request.urlopen(url, timeout=4) as resp:
return resp.read().decode('utf-8', 'ignore')
try:
# thread: the blocking fetch must never stall the event loop
# (an unreachable host parks urlopen for the full 4s timeout)
text = await asyncio.to_thread(_fetch)
except Exception:
self._no_signal_ramp() # no metrics: demand-driven fallback
return
running = queue = None
for line in text.splitlines():
if line.startswith('sglang:num_running_reqs'):
running = float(line.rsplit(' ', 1)[-1])
elif line.startswith('sglang:num_queue_reqs'):
queue = float(line.rsplit(' ', 1)[-1])
self.stats['probe'] += 1
if queue is None and running is None:
self._no_signal_ramp()
return
if queue is not None and queue >= 2:
# server is queuing OUR excess: gentle additive decrease
self.limit = max(self.LO, self.limit - 1)
self.stats['backoff_queue'] += 1
self._push_limit()
elif (queue or 0) == 0 and (running is None or running < max(2, int(self.limit))) \
and self._interval_ok > 0:
# underfed: no queue and running below our own cap -> ramp up
# (still requires completions this interval: no throughput, no ramp)
self.limit = min(self.HI, self.limit + 1)
self.stats['ramp'] += 1
self._push_limit()
else:
self.stats['hold_queue'] += 1
self._interval_fails = 0
self._contended = 0
self._wake()
async def _probe_loop(self) -> None:
import contextlib
while not self._stopped:
with contextlib.suppress(Exception):
await self._probe_once()
await asyncio.sleep(self.PROBE_S)
def stop(self) -> None:
self._stopped = True
if self._task is not None:
self._task.cancel()
def report(self) -> Dict[str, float]:
return {'limit': max(1, int(self.limit)), 'inflight': self._inflight,
**{f'gate_{k}': v for k, v in self.stats.items()}}
def pooled(specs: List[str], api_key: str = '') -> PooledAdapter:
"""['openai/http://127.0.0.1:8123/v1?M', ...] -> PooledAdapter.
api_key: explicit key applied to EVERY member (two-key setups should
build two pools, or use env resolution per host)."""
from .adapter import resolve_adapter
members = [resolve_adapter(s) for s in specs]
if api_key:
for m in members:
m.api_key = api_key
return PooledAdapter(members)