Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 46 additions & 10 deletions src/gabriel/utils/openai_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2567,6 +2567,7 @@ async def worker() -> None:
while True:
if stop_event.is_set():
break
_maybe_emit_ramp_complete()
current_cap = _current_parallel_cap()
if active_workers < current_cap:
break
Expand Down Expand Up @@ -2977,7 +2978,7 @@ async def get_all_responses(
# the API or tool backends.
n_parallels: int = 650,
ramp_up_seconds: float = 20.0,
ramp_up_start_fraction: float = 0.25,
ramp_up_start_fraction: float = 0.15,
max_retries: int = 3,
timeout_factor: float = 2.5,
max_timeout: Optional[float] = None,
Expand Down Expand Up @@ -4348,6 +4349,7 @@ def _effective_parallel_ceiling() -> int:
last_wait_adjust = 0.0
timeout_notes: Deque[str] = deque(maxlen=3)
timeout_errors_since_last_status = 0
rate_limit_errors_since_last_status = 0
connection_errors_since_last_status = 0
first_timeout_logged = False
first_rate_limit_logged = False
Expand Down Expand Up @@ -4379,6 +4381,7 @@ def _effective_parallel_ceiling() -> int:
ramp_up_end_time = ramp_up_start_time + ramp_up_seconds
ramp_up_start_cap = max(1, int(math.ceil(concurrency_cap * ramp_up_start_fraction)))
ramp_up_halted = False
ramp_up_complete_emitted = False

def _maybe_reduce_output_headroom(*, allow_reduce: bool) -> bool:
"""Lower output headroom after the first token estimate refresh."""
Expand Down Expand Up @@ -4474,9 +4477,12 @@ def _log_timeout_once(message: str, note: str, *, dedup_key: Hashable) -> None:
def _log_rate_limit_once(detail: Optional[str] = None) -> None:
nonlocal first_rate_limit_logged
if not first_rate_limit_logged:
logger.warning(
msg = (
"Encountered first rate limit error. Future rate limit errors will be silenced and tracked in periodic updates."
)
logger.warning(msg)
if message_verbose:
print(msg)
first_rate_limit_logged = True
else:
if detail:
Expand All @@ -4487,9 +4493,12 @@ def _log_rate_limit_once(detail: Optional[str] = None) -> None:
def _log_connection_once(detail: Optional[str] = None) -> None:
nonlocal first_connection_logged
if not first_connection_logged:
logger.warning(
msg = (
"Encountered first connection error. Future connection errors will be silenced and tracked in periodic updates."
)
logger.warning(msg)
if message_verbose:
print(msg)
first_connection_logged = True
else:
if detail:
Expand All @@ -4501,6 +4510,8 @@ def _halt_ramp_up(reason: str) -> None:
nonlocal ramp_up_halted, concurrency_cap
if not ramp_up_enabled or ramp_up_halted:
return
if time.time() >= ramp_up_end_time:
return
ramp_cap = _current_ramp_cap()
ramp_up_halted = True
if concurrency_cap > ramp_cap:
Expand Down Expand Up @@ -4553,6 +4564,21 @@ def _trigger_timeout_burst(now: Optional[float] = None) -> None:
if drained and message_verbose:
print(f"[timeouts] Drained {drained} queued prompts before restart.")

def _maybe_emit_ramp_complete(now: Optional[float] = None) -> None:
nonlocal ramp_up_complete_emitted
if not ramp_up_enabled or ramp_up_halted or ramp_up_complete_emitted:
return
now = time.time() if now is None else now
if now < ramp_up_end_time:
return
ramp_up_complete_emitted = True
msg = (
f"[parallelization] Ramp-up complete at {concurrency_cap} parallel threads."
)
logger.info(msg)
if message_verbose:
print(msg)

def _cost_progress_snapshot() -> Optional[Tuple[float, bool]]:
"""Return total cost so far and whether the value is sampled."""

Expand Down Expand Up @@ -4643,9 +4669,9 @@ def emit_parallelization_status(
timeout_text = ""
connection_text = ""
total_completed = processed
denom = max(total_completed, 1)
effective_cap = _current_parallel_cap()
if status.num_timeout_errors or total_completed:
denom = max(total_completed, 1)
timeout_text = f"timeouts={status.num_timeout_errors}/{denom}"
if status_report_interval is not None and timeout_errors_since_last_status:
timeout_text += f" (+{timeout_errors_since_last_status} since last)"
Expand All @@ -4664,17 +4690,21 @@ def emit_parallelization_status(
total_cost, sampled = cost_snapshot
cost_text = f"cost_so_far={'~' if sampled else ''}${total_cost:.2f}"
prefix = f"[{label}] " if label else ""
if status.num_connection_errors or connection_errors_since_last_status:
connection_text = f"connection_errors={status.num_connection_errors}"
if status_report_interval is not None and connection_errors_since_last_status:
connection_text += f" (+{connection_errors_since_last_status} since last)"
connection_text = f"connection_errors={status.num_connection_errors}/{denom}"
if status_report_interval is not None and connection_errors_since_last_status:
connection_text += f" (+{connection_errors_since_last_status} since last)"
rate_limit_text = f"rate_limit_errors={status.num_rate_limit_errors}/{denom}"
if status_report_interval is not None and rate_limit_errors_since_last_status:
rate_limit_text += (
f" (+{rate_limit_errors_since_last_status} since last)"
)
status_bits: List[str] = [
f"cap={effective_cap}",
f"active={active_workers}",
f"inflight={len(inflight)}",
f"queue={queue.qsize()}",
f"processed={processed}/{status.num_tasks_started}",
f"rate_limit_errors={status.num_rate_limit_errors}",
rate_limit_text,
]
if ramp_up_enabled and not ramp_up_halted and effective_cap < concurrency_cap:
status_bits.insert(1, f"ramp_target={concurrency_cap}")
Expand All @@ -4696,7 +4726,7 @@ def emit_parallelization_status(
if ramp_up_enabled and concurrency_cap > 1:
ramp_msg = (
f"[parallelization] Ramping up from {ramp_up_start_cap} to {concurrency_cap} "
f"over {int(ramp_up_seconds)}s."
f"parallel threads over {int(ramp_up_seconds)}s."
)
logger.info(ramp_msg)
if message_verbose:
Expand Down Expand Up @@ -4918,6 +4948,7 @@ def maybe_adjust_concurrency() -> None:
f"after {recent_errors} rate-limit errors in the last {int(round(rate_limit_window))}s."
)
logger.warning(reason)
_halt_ramp_up("rate limit recovery")
emit_parallelization_status(reason, force=True)
else:
concurrency_cap = new_cap
Expand Down Expand Up @@ -4990,6 +5021,7 @@ def maybe_adjust_for_connection_errors() -> None:
"or reduce `n_parallels`."
)
logger.warning(reason)
_halt_ramp_up("connection errors")
emit_parallelization_status(reason, force=True)
connection_errors_since_adjust = 0
connection_error_times.clear()
Expand Down Expand Up @@ -5023,6 +5055,7 @@ async def _maybe_retry(backoff: float) -> bool:

async def _handle_rate_limit_error(error_text: str) -> None:
nonlocal cooldown_until, rate_limit_errors_since_adjust, successes_since_adjust, processed
nonlocal rate_limit_errors_since_last_status
inflight.pop(ident, None)
status.num_rate_limit_errors += 1
status.time_of_last_rate_limit_error = time.time()
Expand All @@ -5031,6 +5064,7 @@ async def _handle_rate_limit_error(error_text: str) -> None:
error_logs[ident].append(error_text)
rate_limit_error_times.append(time.time())
rate_limit_errors_since_adjust += 1
rate_limit_errors_since_last_status += 1
successes_since_adjust = 0
_halt_ramp_up("rate limit error")
if _is_quota_error_message(error_text):
Expand Down Expand Up @@ -5842,6 +5876,7 @@ async def status_reporter() -> None:
return
try:
nonlocal timeout_errors_since_last_status, connection_errors_since_last_status
nonlocal rate_limit_errors_since_last_status
while not stop_event.is_set():
await asyncio.sleep(status_report_interval)
if stop_event.is_set() or processed >= status.num_tasks_started:
Expand All @@ -5850,6 +5885,7 @@ async def status_reporter() -> None:
"Periodic status update", force=True, label=None
)
timeout_errors_since_last_status = 0
rate_limit_errors_since_last_status = 0
connection_errors_since_last_status = 0
except asyncio.CancelledError:
pass
Expand Down
38 changes: 37 additions & 1 deletion tests/test_parallelization_tuning.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ async def responder(prompt: str, **_: object):
assert active["peak"] > 2


def test_ramp_up_halts_on_rate_limit(tmp_path):
def test_ramp_up_halts_on_rate_limit(tmp_path, capsys):
active = {"current": 0, "peak": 0}
lock = asyncio.Lock()
first_error = {"raised": False}
Expand Down Expand Up @@ -168,4 +168,40 @@ async def responder(prompt: str, **_: object):
)
)

output = capsys.readouterr().out
assert "Halting ramp-up" in output
assert active["peak"] <= 3


def test_ramp_up_does_not_halt_after_window(tmp_path, capsys):
first_error = {"raised": False}

async def responder(prompt: str, **_: object):
if not first_error["raised"]:
first_error["raised"] = True
await asyncio.sleep(0.06)
raise openai_utils.RateLimitError("rate limit")
await asyncio.sleep(0.01)
return [f"ok-{prompt}"], 0.01, []

asyncio.run(
openai_utils.get_all_responses(
prompts=[f"p{i}" for i in range(4)],
identifiers=[f"p{i}" for i in range(4)],
response_fn=responder,
use_dummy=False,
save_path=str(tmp_path / "responses.csv"),
reset_files=True,
dynamic_timeout=False,
max_retries=1,
n_parallels=4,
ramp_up_seconds=0.05,
ramp_up_start_fraction=0.25,
status_report_interval=None,
global_cooldown=0,
logging_level="error",
)
)

output = capsys.readouterr().out
assert "Halting ramp-up" not in output
Loading