Spaces:
Running
Running
Download src/tsfm_bench/eval/external_api_predictor.py from ThinkcatLab/LiveHouse-TS: direct link, hf CLI and curl.
- Browser
- Download file 11.6 kB
-
https://huggingface.co/spaces/ThinkcatLab/LiveHouse-TS/resolve/main/src/tsfm_bench/eval/external_api_predictor.py
- Command line
-
hf download hf://spaces/ThinkcatLab/LiveHouse-TS/src/tsfm_bench/eval/external_api_predictor.py
-
curl -L -o external_api_predictor.py https://huggingface.co/spaces/ThinkcatLab/LiveHouse-TS/resolve/main/src/tsfm_bench/eval/external_api_predictor.py
11.6 kB
| from __future__ import annotations | |
| import hashlib | |
| import logging | |
| import os | |
| import time | |
| from dataclasses import dataclass | |
| from typing import Any, Iterator, List, Optional | |
| from urllib.parse import urlparse | |
| import httpx | |
| import numpy as np | |
| from gluonts.dataset import Dataset as GluonDataset | |
| from gluonts.model import Forecast | |
| from gluonts.model.forecast import QuantileForecast | |
| from gluonts.model.predictor import RepresentablePredictor | |
| from tsfm_bench.eval.predictors import _ForecastConfig | |
| logger = logging.getLogger(__name__) | |
| DEFAULT_EXTERNAL_TIMEOUT = 90.0 | |
| DEFAULT_MAX_CONTEXT_POINTS = 4096 | |
| DEFAULT_MAX_RESPONSE_BYTES = 5 * 1024 * 1024 | |
| RETRYABLE_STATUS_CODES = {408, 429, 500, 502, 503, 504} | |
| class ExternalApiConfig: | |
| """Configuration for a user-hosted community forecasting endpoint.""" | |
| endpoint_url: str | |
| model_id: str | |
| auth_token_env: str | None = None | |
| auth_header: str = "Authorization" | |
| timeout: float = DEFAULT_EXTERNAL_TIMEOUT | |
| max_retries: int = 2 | |
| max_context_points: int = DEFAULT_MAX_CONTEXT_POINTS | |
| max_response_bytes: int = DEFAULT_MAX_RESPONSE_BYTES | |
| require_https: bool = True | |
| send_item_metadata: bool = False | |
| class ExternalApiPredictor(RepresentablePredictor): | |
| """Forecast predictor backed by a user-operated HTTPS endpoint. | |
| This adapter intentionally treats the endpoint as a black box. The | |
| leaderboard process never imports user code or downloads user weights; it | |
| only sends the causal context window and validates the returned forecasts. | |
| """ | |
| def __init__( | |
| self, | |
| config: ExternalApiConfig, | |
| prediction_length: int, | |
| quantile_levels: Optional[List[float]] = None, | |
| *, | |
| transport: httpx.BaseTransport | None = None, | |
| ): | |
| super().__init__(prediction_length=prediction_length) | |
| _validate_endpoint_url(config.endpoint_url, require_https=config.require_https) | |
| self.config = config | |
| self.quantile_levels = quantile_levels or [] | |
| self.forecast_config = _ForecastConfig.from_quantiles(self.quantile_levels) | |
| self.leaderboard_name = config.model_id | |
| self._transport = transport | |
| def predict(self, dataset: GluonDataset, **kwargs) -> Iterator[Forecast]: | |
| headers = {"Content-Type": "application/json"} | |
| token = _read_auth_token(self.config.auth_token_env) | |
| if token: | |
| headers[self.config.auth_header] = f"Bearer {token}" | |
| with httpx.Client( | |
| timeout=self.config.timeout, | |
| trust_env=False, | |
| transport=self._transport, | |
| ) as client: | |
| dataset_freq = getattr(dataset, "freq", None) | |
| for entry in dataset: | |
| target = np.asarray(entry["target"], dtype=np.float64) | |
| univariate = target.ndim == 1 | |
| working = target.reshape(1, -1) if univariate else target | |
| freq = _entry_frequency(entry, dataset_freq) | |
| start = _format_start(entry.get("start")) | |
| forecast_rows = [] | |
| for variate_idx in range(working.shape[0]): | |
| series = _finite_tail( | |
| working[variate_idx], | |
| max_points=self.config.max_context_points, | |
| ) | |
| payload = _build_request_payload( | |
| config=self.config, | |
| series=series, | |
| prediction_length=self.prediction_length, | |
| quantile_levels=self.quantile_levels, | |
| freq=freq, | |
| start=start, | |
| item_id=str(entry.get("item_id", variate_idx)), | |
| variate_idx=variate_idx, | |
| ) | |
| response_payload = _post_with_retries( | |
| client, | |
| self.config, | |
| headers=headers, | |
| payload=payload, | |
| ) | |
| forecast_rows.append( | |
| parse_external_forecast( | |
| response_payload, | |
| prediction_length=self.prediction_length, | |
| quantile_levels=self.quantile_levels, | |
| ) | |
| ) | |
| stacked = np.stack(forecast_rows, axis=0) | |
| forecast_arrays = stacked[0] if univariate else stacked | |
| yield QuantileForecast( | |
| forecast_arrays=forecast_arrays, | |
| forecast_keys=self.forecast_config.forecast_keys, | |
| start_date=entry["start"] + target.shape[-1], | |
| item_id=entry["item_id"], | |
| ) | |
| def _validate_endpoint_url(endpoint_url: str, *, require_https: bool = True) -> None: | |
| parsed = urlparse(endpoint_url) | |
| if parsed.scheme not in {"http", "https"}: | |
| raise ValueError("External model endpoint must use http:// or https://") | |
| if not parsed.netloc: | |
| raise ValueError("External model endpoint URL is missing a host") | |
| host = parsed.hostname or "" | |
| local_host = host in {"localhost", "127.0.0.1", "::1"} | |
| if require_https and parsed.scheme != "https" and not local_host: | |
| raise ValueError("External model endpoint must use HTTPS") | |
| def _read_auth_token(env_name: str | None) -> str | None: | |
| if not env_name: | |
| return None | |
| token = os.getenv(env_name) | |
| return token.strip() if token else None | |
| def _entry_frequency(entry: dict[str, Any], dataset_freq: object | None) -> str: | |
| if dataset_freq: | |
| return str(dataset_freq) | |
| if entry.get("freq"): | |
| return str(entry["freq"]) | |
| start = entry.get("start") | |
| return str(getattr(start, "freqstr", None) or "H") | |
| def _format_start(start: Any) -> str: | |
| if hasattr(start, "to_timestamp"): | |
| return start.to_timestamp().isoformat() | |
| return str(start) | |
| def _finite_tail(series: np.ndarray, *, max_points: int) -> np.ndarray: | |
| values = np.asarray(series, dtype=np.float64).reshape(-1) | |
| values = values[np.isfinite(values)] | |
| if values.size == 0: | |
| values = np.zeros(1, dtype=np.float64) | |
| if values.size > max_points: | |
| values = values[-max_points:] | |
| return values | |
| def _opaque_item_id(item_id: str, variate_idx: int) -> str: | |
| digest = hashlib.sha256(f"{item_id}:{variate_idx}".encode("utf-8")).hexdigest()[:16] | |
| return f"series-{digest}" | |
| def _build_request_payload( | |
| *, | |
| config: ExternalApiConfig, | |
| series: np.ndarray, | |
| prediction_length: int, | |
| quantile_levels: list[float], | |
| freq: str, | |
| start: str, | |
| item_id: str, | |
| variate_idx: int, | |
| ) -> dict[str, Any]: | |
| input_block: dict[str, Any] = { | |
| "series_id": _opaque_item_id(item_id, variate_idx), | |
| "target": [float(value) for value in series.tolist()], | |
| } | |
| if config.send_item_metadata: | |
| input_block.update({"item_id": item_id, "start": start, "freq": freq}) | |
| return { | |
| "protocol_version": "tsfm-realworld-v1", | |
| "model": config.model_id, | |
| "inputs": [input_block], | |
| "parameters": { | |
| "prediction_length": int(prediction_length), | |
| "freq": freq, | |
| "quantiles": [float(level) for level in quantile_levels], | |
| }, | |
| } | |
| def _post_with_retries( | |
| client: httpx.Client, | |
| config: ExternalApiConfig, | |
| *, | |
| headers: dict[str, str], | |
| payload: dict[str, Any], | |
| ) -> dict[str, Any]: | |
| attempts = max(1, int(config.max_retries) + 1) | |
| last_error: Exception | None = None | |
| for attempt in range(1, attempts + 1): | |
| try: | |
| response = client.post(config.endpoint_url, headers=headers, json=payload) | |
| if response.status_code in RETRYABLE_STATUS_CODES and attempt < attempts: | |
| _sleep_before_retry(attempt) | |
| continue | |
| response.raise_for_status() | |
| if len(response.content) > config.max_response_bytes: | |
| raise ValueError( | |
| f"External endpoint response is too large: " | |
| f"{len(response.content)} > {config.max_response_bytes} bytes" | |
| ) | |
| data = response.json() | |
| if not isinstance(data, dict): | |
| raise ValueError("External endpoint response must be a JSON object") | |
| return data | |
| except (httpx.HTTPError, ValueError) as exc: | |
| last_error = exc | |
| if attempt >= attempts: | |
| break | |
| logger.warning( | |
| "External endpoint request failed for %s: %s; retrying (%d/%d)", | |
| config.model_id, | |
| exc, | |
| attempt, | |
| attempts, | |
| ) | |
| _sleep_before_retry(attempt) | |
| raise RuntimeError(f"External endpoint failed for {config.model_id}: {last_error}") | |
| def _sleep_before_retry(attempt: int) -> None: | |
| time.sleep(min(8, 2 ** attempt)) | |
| def parse_external_forecast( | |
| payload: dict[str, Any], | |
| *, | |
| prediction_length: int, | |
| quantile_levels: list[float], | |
| ) -> np.ndarray: | |
| """Parse the community endpoint forecast response into GluonTS arrays.""" | |
| outputs = payload.get("outputs") or payload.get("forecasts") | |
| if not isinstance(outputs, list) or not outputs: | |
| raise ValueError("External endpoint response must include a non-empty outputs list") | |
| first = outputs[0] | |
| if not isinstance(first, dict): | |
| raise ValueError("External endpoint output item must be an object") | |
| mean_values = first.get("mean") | |
| if mean_values is None: | |
| mean_values = first.get("prediction") | |
| if mean_values is None: | |
| mean_values = _lookup_quantile(first, 0.5) | |
| if mean_values is None: | |
| raise ValueError("External endpoint response is missing mean or median forecast") | |
| rows = [_forecast_vector(mean_values, prediction_length)] | |
| for level in quantile_levels: | |
| q_values = _lookup_quantile(first, level) | |
| rows.append(rows[0].copy() if q_values is None else _forecast_vector(q_values, prediction_length)) | |
| forecast = np.stack(rows, axis=0).astype(np.float64) | |
| if not np.all(np.isfinite(forecast)): | |
| raise ValueError("External endpoint returned non-finite forecast values") | |
| return forecast | |
| def _lookup_quantile(output: dict[str, Any], level: float) -> Any | None: | |
| for container_name in ("quantiles", "quantile_predictions"): | |
| container = output.get(container_name) | |
| if isinstance(container, dict): | |
| for key in (f"{level:g}", str(level), f"q{level:g}"): | |
| if key in container: | |
| return container[key] | |
| elif isinstance(container, list): | |
| for item in container: | |
| if not isinstance(item, dict): | |
| continue | |
| raw_level = item.get("level", item.get("quantile")) | |
| try: | |
| item_level = float(raw_level) | |
| except (TypeError, ValueError): | |
| continue | |
| if abs(item_level - level) < 1e-8: | |
| return item.get("values") | |
| return None | |
| def _forecast_vector(values: Any, prediction_length: int) -> np.ndarray: | |
| arr = np.asarray(values, dtype=np.float64) | |
| if arr.ndim == 2 and arr.shape[1] == 1: | |
| arr = arr[:, 0] | |
| elif arr.ndim > 1: | |
| arr = arr.reshape(arr.shape[0], -1).mean(axis=1) | |
| arr = arr.reshape(-1) | |
| if arr.size < prediction_length: | |
| raise ValueError(f"Forecast length {arr.size} < expected {prediction_length}") | |
| return arr[:prediction_length] | |