stevenmcdermott commited on
Commit
1778581
Β·
verified Β·
1 Parent(s): db94bdf

Upload api/app.py

Browse files
Files changed (1) hide show
  1. api/app.py +275 -0
api/app.py ADDED
@@ -0,0 +1,275 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ api/app.py β€” FastAPI backend for the TS Anomaly Detection Benchmark
3
+ ====================================================================
4
+
5
+ Designed to run on Hugging Face Spaces (Docker, port 7860).
6
+
7
+ Endpoints:
8
+ POST /api/run β€” start a benchmark job
9
+ GET /api/status/{job_id} β€” poll for progress + results
10
+ GET / β€” health check
11
+ """
12
+
13
+ import os
14
+ import sys
15
+ import uuid
16
+ import math
17
+ import base64
18
+ import threading
19
+ import io
20
+ from typing import Optional, List
21
+
22
+ from fastapi import FastAPI, HTTPException
23
+ from fastapi.middleware.cors import CORSMiddleware
24
+ from pydantic import BaseModel
25
+
26
+ # ── Add benchmark package to Python path ──
27
+ # Allows importing data.synthetic, evaluation.scorer, models.*, etc.
28
+ _BENCH_PATH = os.path.join(os.path.dirname(__file__), "..", "ts-anomaly-benchmark")
29
+ sys.path.insert(0, os.path.abspath(_BENCH_PATH))
30
+
31
+ # ── App ──────────────────────────────────────────────────────────────
32
+ app = FastAPI(title="TS Anomaly Benchmark API", version="1.0.0")
33
+
34
+ app.add_middleware(
35
+ CORSMiddleware,
36
+ allow_origins=["*"],
37
+ allow_methods=["*"],
38
+ allow_headers=["*"],
39
+ )
40
+
41
+ # ── Job store ─────────────────────────────────────────────────────────
42
+ # Simple in-memory store. Fine for a demo β€” one server instance.
43
+ _jobs: dict = {}
44
+ _run_lock = threading.Lock() # one benchmark at a time (models are CPU-heavy)
45
+
46
+
47
+ # ── Request / Response models ─────────────────────────────────────────
48
+
49
+ class RunRequest(BaseModel):
50
+ models: List[str] # e.g. ["moment", "isolation_forest"]
51
+ synthetic_types: List[str] = [
52
+ "sine_with_spikes",
53
+ "random_walk_with_shift",
54
+ "seasonal_with_noise",
55
+ "ecg_like",
56
+ ]
57
+ num_points: int = 512
58
+ custom_csv: Optional[str] = None # base64-encoded CSV content
59
+ custom_value_col: str = "value"
60
+ custom_label_col: Optional[str] = None
61
+
62
+
63
+ # ── Endpoints ─────────────────────────────────────────────────────────
64
+
65
+ @app.get("/")
66
+ def root():
67
+ return {"status": "ok", "message": "TS Anomaly Benchmark API"}
68
+
69
+
70
+ @app.post("/api/run")
71
+ def start_run(req: RunRequest):
72
+ """Start a benchmark job. Returns a job_id for polling."""
73
+ job_id = str(uuid.uuid4())
74
+ _jobs[job_id] = {
75
+ "status": "running",
76
+ "logs": [],
77
+ "results": None,
78
+ "error": None,
79
+ }
80
+ thread = threading.Thread(
81
+ target=_execute_job,
82
+ args=(job_id, req),
83
+ daemon=True,
84
+ )
85
+ thread.start()
86
+ return {"job_id": job_id}
87
+
88
+
89
+ @app.get("/api/status/{job_id}")
90
+ def get_status(job_id: str):
91
+ """Poll job progress. When status=='done', results are included."""
92
+ job = _jobs.get(job_id)
93
+ if job is None:
94
+ raise HTTPException(status_code=404, detail="Job not found")
95
+ return {
96
+ "status": job["status"],
97
+ "logs": job["logs"],
98
+ "results": job["results"],
99
+ "error": job["error"],
100
+ }
101
+
102
+
103
+ # ── Job execution ─────────────────────────────────────────────────────
104
+
105
+ def _execute_job(job_id: str, req: RunRequest):
106
+ """Run the full benchmark in a background thread."""
107
+ job = _jobs[job_id]
108
+ logs = job["logs"]
109
+
110
+ # Serialise concurrent jobs β€” models are memory-heavy
111
+ with _run_lock:
112
+ original_stdout = sys.stdout
113
+ sys.stdout = _LogCapture(logs, original_stdout)
114
+ try:
115
+ config = _build_config(req)
116
+
117
+ # Import benchmark modules (torch may be slow to first import)
118
+ from data.synthetic import generate_all as generate_synthetic
119
+ from evaluation.scorer import build_models, run_benchmark
120
+
121
+ # ── Datasets ──
122
+ all_datasets = {}
123
+ if req.synthetic_types:
124
+ all_datasets.update(generate_synthetic(config["datasets"]["synthetic"]))
125
+ if req.custom_csv:
126
+ all_datasets.update(_load_custom_csv(req))
127
+
128
+ if not all_datasets:
129
+ raise ValueError("No datasets loaded.")
130
+
131
+ # ── Models ──
132
+ models = build_models(config["models"])
133
+ if not models:
134
+ raise ValueError("No models enabled. Select at least one.")
135
+
136
+ # ── Run ──
137
+ results_df = run_benchmark(models, all_datasets, config["evaluation"])
138
+
139
+ job["results"] = _serialise(results_df, all_datasets)
140
+ job["status"] = "done"
141
+
142
+ except Exception as exc:
143
+ job["status"] = "error"
144
+ job["error"] = str(exc)
145
+ logs.append(f"ERROR: {exc}")
146
+ finally:
147
+ sys.stdout = original_stdout
148
+
149
+
150
+ def _build_config(req: RunRequest) -> dict:
151
+ enabled = set(req.models)
152
+ return {
153
+ "models": {
154
+ "moment": {
155
+ "enabled": "moment" in enabled,
156
+ "pretrained": "AutonLab/MOMENT-1-large",
157
+ "task": "reconstruction",
158
+ },
159
+ "isolation_forest": {
160
+ "enabled": "isolation_forest" in enabled,
161
+ "n_estimators": 100,
162
+ "contamination": "auto",
163
+ "window_features": True,
164
+ "feature_window": 20,
165
+ },
166
+ "lof": {
167
+ "enabled": "lof" in enabled,
168
+ "n_neighbors": 20,
169
+ "contamination": "auto",
170
+ "window_features": True,
171
+ "feature_window": 20,
172
+ },
173
+ "moving_window": {
174
+ "enabled": "moving_window" in enabled,
175
+ "window_size": 15,
176
+ "sigma_multiplier": 2.0,
177
+ },
178
+ },
179
+ "datasets": {
180
+ "synthetic": {
181
+ "enabled": bool(req.synthetic_types),
182
+ "num_points": req.num_points,
183
+ "types": req.synthetic_types,
184
+ "anomaly_rate": 0.03,
185
+ "random_seed": 42,
186
+ }
187
+ },
188
+ "evaluation": {"threshold_method": "best_f1"},
189
+ }
190
+
191
+
192
+ def _load_custom_csv(req: RunRequest) -> dict:
193
+ import pandas as pd
194
+ import numpy as np
195
+
196
+ csv_bytes = base64.b64decode(req.custom_csv)
197
+ df = pd.read_csv(io.BytesIO(csv_bytes))
198
+
199
+ if req.custom_value_col not in df.columns:
200
+ numeric = df.select_dtypes(include=[np.number]).columns
201
+ if len(numeric) == 0:
202
+ raise ValueError("No numeric columns found in uploaded CSV.")
203
+ value_col = numeric[0]
204
+ else:
205
+ value_col = req.custom_value_col
206
+
207
+ series = df[value_col].to_numpy(dtype=float)
208
+ labels = None
209
+ if req.custom_label_col and req.custom_label_col in df.columns:
210
+ labels = df[req.custom_label_col].to_numpy(dtype=int)
211
+
212
+ valid = ~__import__("numpy").isnan(series)
213
+ series = series[valid]
214
+ if labels is not None:
215
+ labels = labels[valid]
216
+
217
+ return {"custom_upload": {"series": series, "labels": labels}}
218
+
219
+
220
+ def _serialise(results_df, all_datasets: dict) -> dict:
221
+ metrics = []
222
+ for _, row in results_df.iterrows():
223
+ metrics.append({
224
+ "model": row.get("model"),
225
+ "dataset": row.get("dataset"),
226
+ "auc_roc": _f(row.get("auc_roc")),
227
+ "auc_pr": _f(row.get("auc_pr")),
228
+ "f1": _f(row.get("f1")),
229
+ "precision": _f(row.get("precision")),
230
+ "recall": _f(row.get("recall")),
231
+ "time_seconds": _f(row.get("time_seconds")),
232
+ "scores": _arr(row.get("_scores")),
233
+ })
234
+
235
+ series_out = {}
236
+ for name, data in all_datasets.items():
237
+ series_out[name] = {
238
+ "series": [round(float(v), 4) for v in data["series"]],
239
+ "labels": data["labels"].tolist() if data.get("labels") is not None else None,
240
+ }
241
+
242
+ return {"metrics": metrics, "series": series_out}
243
+
244
+
245
+ def _f(v):
246
+ try:
247
+ f = float(v)
248
+ return None if math.isnan(f) else round(f, 4)
249
+ except Exception:
250
+ return None
251
+
252
+
253
+ def _arr(v):
254
+ try:
255
+ if v is None or not hasattr(v, "__len__"):
256
+ return None
257
+ return [round(float(x), 4) for x in v]
258
+ except Exception:
259
+ return None
260
+
261
+
262
+ # ── Stdout capture ────────────────────────────────────────────────────
263
+
264
+ class _LogCapture:
265
+ def __init__(self, log_list: list, original):
266
+ self._log = log_list
267
+ self._orig = original
268
+
269
+ def write(self, text: str):
270
+ if text and text.strip():
271
+ self._log.append(text.strip())
272
+ self._orig.write(text)
273
+
274
+ def flush(self):
275
+ self._orig.flush()