Исходный код app.predict.predictor

"""The predictor: every tick, finds each vehicle's target stop 10–15 minutes ahead and gets a delay forecast.

Target = the first planned visit in (T + 10 min, T + 15 min] — the dataset's own definition, so the
horizon criterion holds by construction. One prediction per (vehicle, target stop): as T moves on, the
target rolls forward and a fresh prediction follows, roughly every 1–3 minutes per vehicle.

If the ML service is unreachable, the backend issues its own baseline forecast (delay = current
deviation, risk from the same sigmoid the ML contract uses) marked ``source="fallback"`` and retries ML
after ``retry_s`` — the dashboard keeps working in degraded mode (criterion 5).
"""

from __future__ import annotations

import asyncio
import logging
import math
import time
from collections import deque
from collections.abc import Callable, Iterator
from dataclasses import asdict, dataclass, field
from datetime import datetime
from typing import Literal

from ..clock import DatasetClock
from ..state import Fleet, StopVisit, VehicleState
from ..state.geo import epoch_s
from .client import MlClient, MlUnavailable
from .payload import predict_request, schedule_rows

log = logging.getLogger(__name__)

HORIZON_S = (600, 900)
RISK_RED, RISK_YELLOW = 0.7, 0.35  # same thresholds as the dashboard (frontend/js/config.js)


[документация] @dataclass(frozen=True, slots=True) class Prediction: tr_id: int sample_id: str T: datetime # dataset time the forecast was made at target_pos: int target_stop_id: int target_plan: datetime cur_dev_s: float | None # current deviation sent to the model (None: unknown, 0 was sent) delay_pred_s: float risk_score: float risk_level: str confidence: float source: Literal["ml", "fallback"] made_at: float # wall clock response: dict = field(repr=False) # the full ML response (or the fallback's equivalent) @property def lead_s(self) -> float: return (self.target_plan - self.T).total_seconds() @property def horizon_ok(self) -> bool: return HORIZON_S[0] < self.lead_s <= HORIZON_S[1]
[документация] @dataclass(slots=True) class PredictorStats: batches: int = 0 predictions_ml: int = 0 predictions_fallback: int = 0 ml_errors: int = 0 horizon_ok: int = 0
[документация] class Predictor:
[документация] def __init__(self, fleet: Fleet, clock: DatasetClock, client: MlClient, *, max_ping_age_s: float = 900.0, telemetry_span_s: float = 4500.0, batch_max: int = 64, retry_s: float = 30.0, wall: Callable[[], float] = time.time) -> None: self.fleet = fleet self.clock = clock self.client = client self.max_ping_age_s = max_ping_age_s self.telemetry_span_s = telemetry_span_s self.batch_max = batch_max self.retry_s = retry_s self.stats = PredictorStats() self.latest: dict[int, Prediction] = {} self.recent: deque[Prediction] = deque(maxlen=5000) self.ml_available: bool | None = None # None: not tried yet self.last_error: str | None = None self.ml_batch_ms: deque[float] = deque(maxlen=500) self.last_requests: dict[int, dict] = {} # tr_id → last request body (reused by what-if) self._done: dict[tuple[int, int], str] = {} # (tr_id, target stop) → source of its prediction self._schedule_rows: dict[int, list[dict]] = {} self._retry_at = 0.0 self._listeners: list[Callable[[Prediction], None]] = [] self._wall = wall
[документация] def on_prediction(self, fn: Callable[[Prediction], None]) -> None: self._listeners.append(fn)
# ------------------------------------------------------------------ what's due
[документация] def due(self, T: datetime) -> Iterator[tuple[VehicleState, StopVisit]]: t_s = math.floor(epoch_s(T)) ml_retry_open = self._wall() >= self._retry_at for v in self.fleet.vehicles.values(): s, last = v.schedule, v.last_ping if s is None or last is None or (T - last.event_time).total_seconds() > self.max_ping_age_s: continue # not scheduled, or not in service right now i = s.first_after(t_s + HORIZON_S[0]) if i is None or s.plan_s[i] > t_s + HORIZON_S[1]: continue # no stop planned in the window target = s.visits[i] done = self._done.get((v.tr_id, target.stop_id)) if done == "ml" or (done == "fallback" and not ml_retry_open): continue yield v, target
# ------------------------------------------------------------------ one round
[документация] async def tick(self, T: datetime | None = None) -> int: T = self.clock.now() if T is None else T due = list(self.due(T)) for i in range(0, len(due), self.batch_max): await self._predict(T, due[i:i + self.batch_max]) return len(due)
async def _predict(self, T: datetime, chunk: list[tuple[VehicleState, StopVisit]]) -> None: reqs = [] for v, target in chunk: rows = self._schedule_rows.get(v.tr_id) if rows is None: rows = self._schedule_rows[v.tr_id] = schedule_rows(v.schedule) cur = v.derived.cur_dev_s if v.derived is not None else None req = predict_request(v.tr_id, T, target, 0.0 if cur is None else cur, self.fleet.telemetry(v.tr_id, T, span_s=self.telemetry_span_s), rows) reqs.append(req) self.last_requests[v.tr_id] = req source: Literal["ml", "fallback"] = "fallback" responses = None if self._wall() >= self._retry_at: started = time.perf_counter() try: responses = await self.client.predict_batch(reqs) self.ml_batch_ms.append((time.perf_counter() - started) * 1000) source = "ml" if self.ml_available is not True: log.info("ML service available at %s", self.client.base_url) self.ml_available, self.last_error = True, None except MlUnavailable as e: self.stats.ml_errors += 1 if self.ml_available is not False: log.warning("ML service unavailable, using fallback forecasts: %s", e) self.ml_available, self.last_error = False, str(e) self._retry_at = self._wall() + self.retry_s if responses is None: responses = [fallback_response(r) for r in reqs] self.stats.batches += 1 now = self._wall() for (v, target), req, resp in zip(chunk, reqs, responses): cur = v.derived.cur_dev_s if v.derived is not None else None p = Prediction(tr_id=v.tr_id, sample_id=req["sample_id"], T=T, target_pos=target.pos, target_stop_id=target.stop_id, target_plan=target.plan, cur_dev_s=cur, delay_pred_s=float(resp["delay_pred_sec"]), risk_score=float(resp["risk_score"]), risk_level=resp.get("risk_level") or level_for(float(resp["risk_score"])), confidence=float(resp.get("confidence", 0.0)), source=source, made_at=now, response=resp) self._record(p) def _record(self, p: Prediction) -> None: self._done[(p.tr_id, p.target_stop_id)] = p.source self.latest[p.tr_id] = p self.recent.append(p) if p.source == "ml": self.stats.predictions_ml += 1 else: self.stats.predictions_fallback += 1 self.stats.horizon_ok += p.horizon_ok for fn in self._listeners: try: fn(p) except Exception: # noqa: BLE001 — a broken consumer must not stop predictions log.exception("prediction listener %r failed", fn)
[документация] async def run(self, tick_s: float = 5.0) -> None: while True: try: await self.tick() except Exception: # noqa: BLE001 — keep predicting on the next tick log.exception("predictor tick failed") await asyncio.sleep(tick_s)
# ------------------------------------------------------------------ introspection
[документация] def snapshot(self) -> dict: total = self.stats.predictions_ml + self.stats.predictions_fallback ms = sorted(self.ml_batch_ms) return { "ml_url": self.client.base_url, "ml_available": self.ml_available, "last_error": self.last_error, "stats": asdict(self.stats), "vehicles_with_prediction": len(self.latest), "horizon_ok_share": round(self.stats.horizon_ok / total, 3) if total else None, "ml_batch_ms_p50": round(ms[len(ms) // 2], 1) if ms else None, "ml_batch_ms_p95": round(ms[min(len(ms) - 1, int(len(ms) * 0.95))], 1) if ms else None, }
[документация] def level_for(risk: float) -> str: return "red" if risk >= RISK_RED else "yellow" if risk >= RISK_YELLOW else "green"
[документация] def risk_from_delay(d: float) -> float: """The ML contract's fallback risk: sigmoid((delay − 120) / 60), i.e. 0.5 at +2 min late.""" z = (d - 120.0) / 60.0 return 1 / (1 + math.exp(-z)) if z >= 0 else math.exp(z) / (1 + math.exp(z))
[документация] def fallback_response(req: dict) -> dict: """What the backend answers itself when ML is down: the baseline (delay = current deviation), same shape as ML.""" d = float(req["cur_dev_s"]) risk = risk_from_delay(d) level = level_for(risk) return { "sample_id": req["sample_id"], "delay_pred_sec": round(d, 1), "risk_score": round(risk, 3), "confidence": 0.05, "top_features": [{"name": "cur_dev_s", "value": d, "contribution": 1.0, "contribution_sec": d}], "reason_pattern": "accumulated_delay" if level != "green" else "on_track", "recommendation": "release_reserve" if level == "red" else "monitor", "recommendation_text": "", "model_version": "backend-fallback", "risk_level": level, "delay_interval_sec": [d - 150, d + 150], "data_status": "fallback", "causes": [{"code": "ml_unavailable", "text": "ML-сервис недоступен: прогноз по текущему отклонению", "contribution_sec": d}], }