File size: 4,695 Bytes
e201ae4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | """ml.executor — bounded, timed offload for the synchronous ML fits (analytical tools round).
The analytical tools (analyze_forecast / importance / anomaly / cluster /
significance) call blocking, CPU-heavy fit routines (scikit-learn, statsmodels).
Their compute functions are SYNCHRONOUS by design — like every other `analyze_*`
fn (see the STATUS note in `analytics/temporal.py`) — but the invoker runs inside
the async slow path. Calling a multi-second fit directly on the event loop would
stall every other request in the process.
This offloads a fit to a DEDICATED, bounded thread pool with a per-fit timeout —
the same pattern the DB executor uses (`src/query/executor/db.py`, F-5):
* Dedicated pool — a slow/hung fit occupies an ML-only thread, so its blast
radius is other analytics, not the whole service (the shared default
`asyncio.to_thread` pool also serves the tabular Parquet loader).
* Semaphore — bounds how many fits run at once, so a burst of heavy questions
QUEUES instead of oversubscribing the CPU (kept equal to the pool size).
* Per-fit timeout — `asyncio.wait_for` stops US waiting. NOTE (the db.py
caveat): it cancels the awaiting COROUTINE, not the worker thread, which is
not cancellable and runs to completion; the pool bound caps how many such
abandoned threads can pile up. sklearn / statsmodels release the GIL inside
their C/Fortran kernels, so threads still give real parallelism here.
The invoker's never-throw seam (§8.4) catches whatever this raises — including
`MLFitTimeoutError` — and returns a `ToolOutput(kind="error")`, so a slow fit
degrades that one analysis, never the whole turn.
"""
from __future__ import annotations
import asyncio
import functools
import os
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from typing import Any, TypeVar
from src.middlewares.logging import get_logger
logger = get_logger("ml_executor")
_T = TypeVar("_T")
def _int_env(name: str, default: int) -> int:
"""A positive int from the environment, falling back to `default` if unset/bad."""
raw = os.getenv(name)
if raw is None:
return default
try:
return max(1, int(raw))
except ValueError:
logger.warning("ignoring non-integer env override", var=name, value=raw)
return default
# ML fits are CPU-bound; more threads than cores just thrash. Small by default,
# overridable for a beefier host (or trimmed on a constrained HF Space).
_MAX_ML_WORKERS = _int_env("ML_MAX_WORKERS", min(4, os.cpu_count() or 1))
_FIT_TIMEOUT_SECONDS = float(_int_env("ML_FIT_TIMEOUT_SECONDS", 30))
_ML_THREAD_POOL = ThreadPoolExecutor(
max_workers=_MAX_ML_WORKERS, thread_name_prefix="mlfit"
)
# Equal to the pool size: excess fits wait here rather than piling futures onto a
# saturated pool. Constructed without a running loop (fine on 3.10+: asyncio
# primitives bind to the loop lazily, on first await).
_ML_SEMAPHORE = asyncio.Semaphore(_MAX_ML_WORKERS)
class MLFitTimeoutError(TimeoutError):
"""A single ML fit exceeded its per-fit timeout.
Subclasses `TimeoutError` so callers may catch either; the invoker's
never-throw seam turns it into a `ToolOutput(kind="error")`.
"""
async def run_fit(
fn: Callable[..., _T],
*args: Any,
label: str = "",
timeout: float | None = None,
**kwargs: Any,
) -> _T:
"""Run a blocking fit ``fn(*args, **kwargs)`` off the event loop — bounded + timed.
Args:
fn: the synchronous compute/fit callable.
*args, **kwargs: forwarded to `fn` unchanged.
label: names the fit in logs and the timeout message (pass the tool name).
timeout: per-fit budget in seconds; defaults to `_FIT_TIMEOUT_SECONDS`
(env `ML_FIT_TIMEOUT_SECONDS`).
Returns:
Whatever `fn` returns.
Raises:
MLFitTimeoutError: the fit did not finish within `timeout`.
Exception: any error from `fn` is re-raised unchanged (the invoker's
never-throw seam wraps it).
"""
budget = _FIT_TIMEOUT_SECONDS if timeout is None else timeout
loop = asyncio.get_running_loop()
call = functools.partial(fn, *args, **kwargs)
async with _ML_SEMAPHORE:
try:
return await asyncio.wait_for(
loop.run_in_executor(_ML_THREAD_POOL, call), budget
)
except TimeoutError: # asyncio.wait_for raises builtin TimeoutError (3.11+)
name = label or getattr(fn, "__name__", "fit")
logger.warning("ml fit timed out", label=name, timeout=budget)
raise MLFitTimeoutError(f"{name} exceeded {budget:g}s") from None
|