stasaking commited on
Commit
72f4068
·
verified ·
1 Parent(s): 1883c90

Upload main.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. main.py +512 -0
main.py ADDED
@@ -0,0 +1,512 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Blood-Brain Omics Benchmark API
3
+
4
+ FastAPI + DuckDB backend serving benchmark results.
5
+ Loads Parquet files at startup and queries them via DuckDB in-memory.
6
+
7
+ Endpoints:
8
+ GET /api/v1/registry - Full registry metadata
9
+ GET /api/v1/maxn/heatmap - Blood x brain heatmap data
10
+ GET /api/v1/maxn/detail - Per-module results for one combination
11
+ GET /api/v1/h2h/blood - Blood H2H comparison
12
+ GET /api/v1/h2h/blood/summary - H2H win counts
13
+ GET /api/v1/h2h/brain - Brain H2H comparison
14
+ GET /api/v1/h2h/models - Model H2H comparison (future)
15
+ GET /api/v1/temporal - Temporal decay results
16
+ GET /api/v1/features - Feature importance for a target
17
+ GET /api/v1/features/cross - Cross-target feature importance
18
+ """
19
+
20
+ import json
21
+ import os
22
+ from typing import Optional
23
+
24
+ import duckdb
25
+ from fastapi import FastAPI, HTTPException, Query
26
+ from fastapi.middleware.cors import CORSMiddleware
27
+
28
+ # =============================================================================
29
+ # Configuration
30
+ # =============================================================================
31
+
32
+ DATA_DIR = os.environ.get(
33
+ "BENCHMARK_DATA_DIR",
34
+ os.path.join(os.path.dirname(__file__), "..", "data")
35
+ )
36
+
37
+ app = FastAPI(
38
+ title="Blood-Brain Omics Benchmark API",
39
+ version="1.0.0",
40
+ description="Interactive exploration of blood omics → brain phenotype predictions",
41
+ )
42
+
43
+ app.add_middleware(
44
+ CORSMiddleware,
45
+ allow_origins=["*"],
46
+ allow_methods=["*"],
47
+ allow_headers=["*"],
48
+ )
49
+
50
+ # =============================================================================
51
+ # Startup: load data into DuckDB
52
+ # =============================================================================
53
+
54
+ db = None
55
+ registry = None
56
+
57
+
58
+ @app.on_event("startup")
59
+ def startup():
60
+ global db, registry
61
+
62
+ # Load registry
63
+ registry_path = os.path.join(DATA_DIR, "benchmark_registry.json")
64
+ if os.path.exists(registry_path):
65
+ with open(registry_path) as f:
66
+ registry = json.load(f)
67
+ else:
68
+ registry = {"project": {}, "blood_platforms": {}, "brain_targets": {},
69
+ "phases": {}, "models": {}}
70
+
71
+ # Initialize DuckDB
72
+ db = duckdb.connect(":memory:")
73
+
74
+ results_path = os.path.join(DATA_DIR, "benchmark_results.parquet")
75
+ features_path = os.path.join(DATA_DIR, "feature_importance.parquet")
76
+
77
+ if os.path.exists(results_path):
78
+ db.execute(f"""
79
+ CREATE TABLE results AS
80
+ SELECT * FROM read_parquet('{results_path}')
81
+ """)
82
+ n = db.execute("SELECT COUNT(*) FROM results").fetchone()[0]
83
+ print(f"Loaded results: {n} rows")
84
+ else:
85
+ print(f"WARNING: {results_path} not found")
86
+ db.execute("""
87
+ CREATE TABLE results (
88
+ phase VARCHAR, blood_platform VARCHAR, brain_target VARCHAR,
89
+ target VARCHAR, model VARCHAR, h2h_pair VARCHAR,
90
+ temporal_bin VARCHAR, n_samples INT, n_features INT,
91
+ include_covariates BOOLEAN, r2 DOUBLE, pearson DOUBLE,
92
+ mse DOUBLE, mae DOUBLE
93
+ )
94
+ """)
95
+
96
+ if os.path.exists(features_path):
97
+ db.execute(f"""
98
+ CREATE TABLE features AS
99
+ SELECT * FROM read_parquet('{features_path}')
100
+ """)
101
+ n = db.execute("SELECT COUNT(*) FROM features").fetchone()[0]
102
+ print(f"Loaded features: {n} rows")
103
+ else:
104
+ print(f"WARNING: {features_path} not found")
105
+ db.execute("""
106
+ CREATE TABLE features (
107
+ phase VARCHAR, blood_platform VARCHAR, brain_target VARCHAR,
108
+ target VARCHAR, model VARCHAR, feature_name VARCHAR,
109
+ importance DOUBLE, rank SMALLINT
110
+ )
111
+ """)
112
+
113
+
114
+ # =============================================================================
115
+ # Helpers
116
+ # =============================================================================
117
+
118
+ VALID_METRICS = {"r2", "pearson", "mse", "mae"}
119
+
120
+
121
+ def validate_metric(metric: str) -> str:
122
+ if metric not in VALID_METRICS:
123
+ raise HTTPException(400, f"Invalid metric: {metric}. "
124
+ f"Must be one of {VALID_METRICS}")
125
+ return metric
126
+
127
+
128
+ def covariate_phase(base_phase: str, covariates: str) -> str:
129
+ """Map covariates param to actual phase name."""
130
+ if covariates == "none":
131
+ return base_phase
132
+ elif covariates == "only":
133
+ return f"{base_phase}_covonly"
134
+ elif covariates == "included":
135
+ return f"{base_phase}_withcov"
136
+ return base_phase
137
+
138
+
139
+ # =============================================================================
140
+ # Endpoints
141
+ # =============================================================================
142
+
143
+ @app.get("/api/v1/registry")
144
+ def get_registry():
145
+ """Full registry metadata."""
146
+ return registry
147
+
148
+
149
+ @app.get("/api/v1/maxn/heatmap")
150
+ def maxn_heatmap(
151
+ metric: str = Query("r2", description="Metric to aggregate"),
152
+ covariates: str = Query("none", enum=["none", "only", "included"]),
153
+ model: str = Query("TabPFN"),
154
+ ):
155
+ """
156
+ Blood x brain heatmap data.
157
+ Returns median metric across modules for each combination.
158
+ """
159
+ validate_metric(metric)
160
+ phase = covariate_phase("maxn", covariates)
161
+
162
+ rows = db.execute(f"""
163
+ SELECT blood_platform, brain_target,
164
+ MEDIAN({metric}) as value,
165
+ MEDIAN(n_samples) as n_samples,
166
+ COUNT(*) as n_targets
167
+ FROM results
168
+ WHERE phase = ? AND model = ?
169
+ GROUP BY blood_platform, brain_target
170
+ ORDER BY blood_platform, brain_target
171
+ """, [phase, model]).fetchall()
172
+
173
+ if not rows:
174
+ return {"rows": [], "cols": [], "values": [], "n_samples": [],
175
+ "metric": metric, "covariates": covariates}
176
+
177
+ # Build matrix
178
+ blood_set = sorted(set(r[0] for r in rows))
179
+ brain_set = sorted(set(r[1] for r in rows))
180
+
181
+ values = [[None] * len(brain_set) for _ in range(len(blood_set))]
182
+ n_samples = [[None] * len(brain_set) for _ in range(len(blood_set))]
183
+
184
+ blood_idx = {b: i for i, b in enumerate(blood_set)}
185
+ brain_idx = {b: i for i, b in enumerate(brain_set)}
186
+
187
+ for blood, brain, val, ns, nt in rows:
188
+ i = blood_idx[blood]
189
+ j = brain_idx[brain]
190
+ values[i][j] = round(val, 4) if val is not None else None
191
+ n_samples[i][j] = int(ns) if ns is not None else None
192
+
193
+ return {
194
+ "rows": blood_set,
195
+ "cols": brain_set,
196
+ "values": values,
197
+ "n_samples": n_samples,
198
+ "metric": metric,
199
+ "covariates": covariates,
200
+ "model": model,
201
+ }
202
+
203
+
204
+ @app.get("/api/v1/maxn/detail")
205
+ def maxn_detail(
206
+ blood: str = Query(..., description="Blood platform name"),
207
+ brain: str = Query(..., description="Brain target name"),
208
+ covariates: str = Query("none", enum=["none", "only", "included"]),
209
+ model: str = Query("TabPFN"),
210
+ ):
211
+ """Per-module results for one blood x brain combination."""
212
+ phase = covariate_phase("maxn", covariates)
213
+
214
+ rows = db.execute("""
215
+ SELECT target, n_samples, r2, pearson, mse, mae
216
+ FROM results
217
+ WHERE phase = ? AND blood_platform = ? AND brain_target = ?
218
+ AND model = ?
219
+ ORDER BY r2 DESC
220
+ """, [phase, blood, brain, model]).fetchall()
221
+
222
+ return {
223
+ "blood": blood,
224
+ "brain": brain,
225
+ "covariates": covariates,
226
+ "model": model,
227
+ "targets": [
228
+ {"target": r[0], "n_samples": r[1], "r2": r[2],
229
+ "pearson": r[3], "mse": r[4], "mae": r[5]}
230
+ for r in rows
231
+ ],
232
+ }
233
+
234
+
235
+ @app.get("/api/v1/h2h/blood")
236
+ def h2h_blood(
237
+ pair: str = Query(..., description="Platform pair, e.g. SomaScan_vs_TMT"),
238
+ brain: str = Query(..., description="Brain target name"),
239
+ metric: str = Query("r2"),
240
+ covariates: str = Query("none", enum=["none", "included"]),
241
+ model: str = Query("TabPFN"),
242
+ ):
243
+ """Per-module comparison for a blood platform pair on one brain target."""
244
+ validate_metric(metric)
245
+ phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
246
+
247
+ rows = db.execute(f"""
248
+ SELECT target, blood_platform, {metric}
249
+ FROM results
250
+ WHERE phase = ? AND h2h_pair = ? AND brain_target = ? AND model = ?
251
+ ORDER BY target
252
+ """, [phase, pair, brain, model]).fetchall()
253
+
254
+ # Pivot: target -> {platform_a: val, platform_b: val}
255
+ platforms = sorted(set(r[1] for r in rows))
256
+ targets = {}
257
+ for target, platform, val in rows:
258
+ if target not in targets:
259
+ targets[target] = {}
260
+ targets[target][platform] = round(val, 4) if val is not None else None
261
+
262
+ return {
263
+ "pair": pair,
264
+ "brain": brain,
265
+ "platforms": platforms,
266
+ "metric": metric,
267
+ "targets": [
268
+ {"target": t, **vals}
269
+ for t, vals in sorted(targets.items())
270
+ ],
271
+ }
272
+
273
+
274
+ @app.get("/api/v1/h2h/blood/summary")
275
+ def h2h_blood_summary(
276
+ metric: str = Query("r2"),
277
+ covariates: str = Query("none", enum=["none", "included"]),
278
+ model: str = Query("TabPFN"),
279
+ ):
280
+ """Win counts across all blood H2H pairs."""
281
+ validate_metric(metric)
282
+ phase = "h2h_blood_withcov" if covariates == "included" else "h2h_blood"
283
+
284
+ rows = db.execute(f"""
285
+ SELECT h2h_pair, brain_target, target, blood_platform, {metric}
286
+ FROM results
287
+ WHERE phase = ? AND model = ? AND h2h_pair IS NOT NULL
288
+ ORDER BY h2h_pair, brain_target, target
289
+ """, [phase, model]).fetchall()
290
+
291
+ # Count wins per pair
292
+ pair_wins = {}
293
+ current = None
294
+ buffer = {}
295
+
296
+ for pair, brain, target, platform, val in rows:
297
+ key = (pair, brain, target)
298
+ if key != current:
299
+ if current and len(buffer) == 2:
300
+ pair_key = current[0]
301
+ if pair_key not in pair_wins:
302
+ pair_wins[pair_key] = {}
303
+ platforms = list(buffer.keys())
304
+ v0, v1 = buffer[platforms[0]], buffer[platforms[1]]
305
+ if v0 is not None and v1 is not None:
306
+ asc = metric in ("mse", "mae")
307
+ winner = platforms[0] if (v0 < v1 if asc else v0 > v1) else platforms[1]
308
+ pair_wins[pair_key][winner] = pair_wins[pair_key].get(winner, 0) + 1
309
+ current = key
310
+ buffer = {}
311
+ buffer[platform] = val
312
+
313
+ # Process last group
314
+ if current and len(buffer) == 2:
315
+ pair_key = current[0]
316
+ if pair_key not in pair_wins:
317
+ pair_wins[pair_key] = {}
318
+ platforms = list(buffer.keys())
319
+ v0, v1 = buffer[platforms[0]], buffer[platforms[1]]
320
+ if v0 is not None and v1 is not None:
321
+ asc = metric in ("mse", "mae")
322
+ winner = platforms[0] if (v0 < v1 if asc else v0 > v1) else platforms[1]
323
+ pair_wins[pair_key][winner] = pair_wins[pair_key].get(winner, 0) + 1
324
+
325
+ return {
326
+ "metric": metric,
327
+ "covariates": covariates,
328
+ "pairs": pair_wins,
329
+ }
330
+
331
+
332
+ @app.get("/api/v1/h2h/brain")
333
+ def h2h_brain(
334
+ blood: str = Query(..., description="Blood platform name"),
335
+ metric: str = Query("r2"),
336
+ covariates: str = Query("none", enum=["none", "included"]),
337
+ model: str = Query("TabPFN"),
338
+ ):
339
+ """Compare brain targets for one blood platform."""
340
+ validate_metric(metric)
341
+ phase = "h2h_brain_withcov" if covariates == "included" else "h2h_brain"
342
+
343
+ rows = db.execute(f"""
344
+ SELECT h2h_pair, brain_target, target, {metric}
345
+ FROM results
346
+ WHERE phase = ? AND blood_platform = ? AND model = ?
347
+ ORDER BY h2h_pair, target
348
+ """, [phase, blood, model]).fetchall()
349
+
350
+ # Group by pair
351
+ pairs = {}
352
+ for pair, brain, target, val in rows:
353
+ if pair not in pairs:
354
+ pairs[pair] = {}
355
+ if target not in pairs[pair]:
356
+ pairs[pair][target] = {}
357
+ pairs[pair][target][brain] = round(val, 4) if val is not None else None
358
+
359
+ return {
360
+ "blood": blood,
361
+ "metric": metric,
362
+ "pairs": pairs,
363
+ }
364
+
365
+
366
+ @app.get("/api/v1/h2h/models")
367
+ def h2h_models(
368
+ blood: str = Query(...),
369
+ brain: str = Query(...),
370
+ metric: str = Query("r2"),
371
+ covariates: str = Query("none", enum=["none", "included"]),
372
+ ):
373
+ """Compare models for one blood x brain combination (future)."""
374
+ validate_metric(metric)
375
+ phase = covariate_phase("maxn", covariates)
376
+
377
+ rows = db.execute(f"""
378
+ SELECT target, model, {metric}
379
+ FROM results
380
+ WHERE phase = ? AND blood_platform = ? AND brain_target = ?
381
+ ORDER BY target, model
382
+ """, [phase, blood, brain]).fetchall()
383
+
384
+ models = sorted(set(r[1] for r in rows))
385
+ targets = {}
386
+ for target, model, val in rows:
387
+ if target not in targets:
388
+ targets[target] = {}
389
+ targets[target][model] = round(val, 4) if val is not None else None
390
+
391
+ return {
392
+ "blood": blood,
393
+ "brain": brain,
394
+ "models": models,
395
+ "metric": metric,
396
+ "targets": [
397
+ {"target": t, **vals}
398
+ for t, vals in sorted(targets.items())
399
+ ],
400
+ }
401
+
402
+
403
+ @app.get("/api/v1/temporal")
404
+ def temporal(
405
+ target: Optional[str] = Query(None, description="Specific target (e.g., gpath)"),
406
+ metric: str = Query("r2"),
407
+ covariates: str = Query("none", enum=["none", "included"]),
408
+ model: str = Query("TabPFN"),
409
+ ):
410
+ """Temporal decay results: metric across time bins."""
411
+ validate_metric(metric)
412
+ phase = "temporal_withcov" if covariates == "included" else "temporal"
413
+
414
+ if target:
415
+ rows = db.execute(f"""
416
+ SELECT temporal_bin, brain_target, target, {metric}, n_samples
417
+ FROM results
418
+ WHERE phase = ? AND model = ? AND target = ?
419
+ ORDER BY temporal_bin
420
+ """, [phase, model, target]).fetchall()
421
+ else:
422
+ rows = db.execute(f"""
423
+ SELECT temporal_bin, brain_target, target, {metric}, n_samples
424
+ FROM results
425
+ WHERE phase = ? AND model = ?
426
+ ORDER BY temporal_bin, target
427
+ """, [phase, model]).fetchall()
428
+
429
+ results = [
430
+ {"bin": r[0], "brain_target": r[1], "target": r[2],
431
+ "value": round(r[3], 4) if r[3] is not None else None,
432
+ "n_samples": r[4]}
433
+ for r in rows
434
+ ]
435
+
436
+ return {
437
+ "metric": metric,
438
+ "covariates": covariates,
439
+ "results": results,
440
+ }
441
+
442
+
443
+ @app.get("/api/v1/features")
444
+ def feature_importance(
445
+ blood: str = Query(...),
446
+ brain: str = Query(...),
447
+ target: str = Query(...),
448
+ phase: str = Query("maxn"),
449
+ model: str = Query("TabPFN"),
450
+ limit: int = Query(30, ge=1, le=100),
451
+ ):
452
+ """Top features for a specific module/target."""
453
+ rows = db.execute("""
454
+ SELECT feature_name, importance, rank
455
+ FROM features
456
+ WHERE phase = ? AND blood_platform = ? AND brain_target = ?
457
+ AND target = ? AND model = ?
458
+ ORDER BY rank
459
+ LIMIT ?
460
+ """, [phase, blood, brain, target, model, limit]).fetchall()
461
+
462
+ return {
463
+ "blood": blood,
464
+ "brain": brain,
465
+ "target": target,
466
+ "phase": phase,
467
+ "features": [
468
+ {"feature": r[0], "importance": round(r[1], 4), "rank": r[2]}
469
+ for r in rows
470
+ ],
471
+ }
472
+
473
+
474
+ @app.get("/api/v1/features/cross")
475
+ def feature_cross_target(
476
+ blood: str = Query(...),
477
+ brain: str = Query(...),
478
+ phase: str = Query("maxn"),
479
+ model: str = Query("TabPFN"),
480
+ limit: int = Query(30, ge=1, le=100),
481
+ ):
482
+ """Aggregated feature importance across all modules in a brain target."""
483
+ rows = db.execute("""
484
+ SELECT feature_name,
485
+ AVG(importance) as mean_importance,
486
+ COUNT(DISTINCT target) as n_targets,
487
+ MIN(rank) as best_rank
488
+ FROM features
489
+ WHERE phase = ? AND blood_platform = ? AND brain_target = ? AND model = ?
490
+ GROUP BY feature_name
491
+ ORDER BY mean_importance DESC
492
+ LIMIT ?
493
+ """, [phase, blood, brain, model, limit]).fetchall()
494
+
495
+ return {
496
+ "blood": blood,
497
+ "brain": brain,
498
+ "phase": phase,
499
+ "features": [
500
+ {"feature": r[0], "mean_importance": round(r[1], 4),
501
+ "n_targets": r[2], "best_rank": r[3]}
502
+ for r in rows
503
+ ],
504
+ }
505
+
506
+
507
+ @app.get("/api/v1/health")
508
+ def health():
509
+ """Health check."""
510
+ n_results = db.execute("SELECT COUNT(*) FROM results").fetchone()[0]
511
+ n_features = db.execute("SELECT COUNT(*) FROM features").fetchone()[0]
512
+ return {"status": "ok", "n_results": n_results, "n_features": n_features}