Spaces:
Sleeping
Sleeping
| import io | |
| import os | |
| from typing import List | |
| import pickle | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| import torch.nn as nn | |
| from fastapi import FastAPI, File, UploadFile, HTTPException | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from pydantic import BaseModel | |
| from safetensors.torch import load_file | |
| # ----------------------------- | |
| # Config | |
| # ----------------------------- | |
| FEATURES = [ | |
| "ESS_TOTAL", "MCATOT", "GDS_TOTAL", | |
| "MCAALTTM", "MCACUBE", "MCASER7", "MCAABSTR", | |
| "GDSSATIS", "GDSHAPPY", "GDSENRGY", | |
| "ESS1", "ESS2" | |
| ] | |
| TARGETS = [ | |
| "DATSCAN_CAUDATE_R", | |
| "DATSCAN_CAUDATE_L", | |
| "DATSCAN_PUTAMEN_R", | |
| "DATSCAN_PUTAMEN_L" | |
| ] | |
| MODEL_PATHS = { | |
| "state": "model/flake_transformer.safetensors", | |
| "sx": "model/scaler_x.pkl", | |
| "sy": "model/scaler_y.pkl", | |
| } | |
| # ----------------------------- | |
| # App | |
| # ----------------------------- | |
| app = FastAPI(title="PD Biomarker Predictor", version="1.0.0") | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], # Tighten for production | |
| allow_credentials=True, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| # ----------------------------- | |
| # Model definition | |
| # ----------------------------- | |
| class flakeParkinsonTransformer(nn.Module): | |
| def __init__(self, num_features: int, output_dim: int, embed_dim=64, nhead=4, num_layers=3): | |
| super().__init__() | |
| self.feature_embeds = nn.ModuleList([nn.Linear(1, embed_dim) for _ in range(num_features)]) | |
| encoder_layer = nn.TransformerEncoderLayer( | |
| d_model=embed_dim, nhead=nhead, dim_feedforward=256, dropout=0.1 | |
| ) | |
| self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) | |
| self.pool = nn.AdaptiveAvgPool1d(1) | |
| self.fc = nn.Sequential( | |
| nn.Linear(embed_dim, 128), | |
| nn.ReLU(), | |
| nn.Dropout(0.1), | |
| nn.Linear(128, output_dim), | |
| ) | |
| def forward(self, x): | |
| embeds = [self.feature_embeds[i](x[:, i].unsqueeze(1)) for i in range(x.shape[1])] | |
| x = torch.stack(embeds, dim=1) | |
| x = x.permute(1, 0, 2) | |
| x = self.transformer(x) | |
| x = x.permute(1, 2, 0) | |
| x = self.pool(x).squeeze(2) | |
| x = self.fc(x) | |
| return x | |
| # ----------------------------- | |
| # Load model and scalers | |
| # ----------------------------- | |
| def load_model_and_scalers(): | |
| # Load scalers using pickle | |
| try: | |
| with open(MODEL_PATHS["sx"], 'rb') as f: | |
| scaler_x = pickle.load(f) | |
| print("scaler_x loaded successfully.") | |
| except Exception as e: | |
| raise RuntimeError(f"Error loading scaler_x: {e}") | |
| try: | |
| with open(MODEL_PATHS["sy"], 'rb') as f: | |
| scaler_y = pickle.load(f) | |
| print("scaler_y loaded successfully.") | |
| except Exception as e: | |
| raise RuntimeError(f"Error loading scaler_y: {e}") | |
| # Load model using safetensors | |
| try: | |
| state_dict = load_file(MODEL_PATHS["state"], device="cpu") | |
| model = flakeParkinsonTransformer(num_features=len(FEATURES), output_dim=len(TARGETS)) | |
| model.load_state_dict(state_dict) | |
| model.eval() | |
| print("Model loaded and set to evaluation mode.") | |
| except Exception as e: | |
| raise RuntimeError(f"Error loading model state dictionary: {e}") | |
| return model, scaler_x, scaler_y | |
| MODEL, SCALER_X, SCALER_Y = load_model_and_scalers() | |
| # ----------------------------- | |
| # Utilities | |
| # ----------------------------- | |
| def try_compute_ess_total(df: pd.DataFrame) -> pd.DataFrame: | |
| if "ESS_TOTAL" not in df.columns: | |
| ess_cols = [f"ESS{i}" for i in range(1, 9) if f"ESS{i}" in df.columns] | |
| if len(ess_cols) == 8: | |
| df["ESS_TOTAL"] = df[ess_cols].sum(axis=1) | |
| return df | |
| def try_compute_gds_total(df: pd.DataFrame) -> pd.DataFrame: | |
| gds_items = [ | |
| "GDSSATIS", "GDSDROPD", "GDSEMPTY", "GDSBORED", "GDSGSPIR", | |
| "GDSAFRAD", "GDSHAPPY", "GDSHLPLS", "GDSHOME", "GDSMEMRY", | |
| "GDSALIVE", "GDSWRTLS", "GDSENRGY", "GDSHOPLS", "GDSBETER" | |
| ] | |
| if "GDS_TOTAL" not in df.columns: | |
| present = [c for c in gds_items if c in df.columns] | |
| if len(present) == 15: | |
| df["GDS_TOTAL"] = df[present].sum(axis=1) | |
| return df | |
| def standardize_and_predict(df_features: pd.DataFrame): | |
| X = df_features[FEATURES].values | |
| X = SCALER_X.transform(X) | |
| X_t = torch.tensor(X, dtype=torch.float32) | |
| with torch.no_grad(): | |
| preds = MODEL(X_t).numpy() | |
| preds = SCALER_Y.inverse_transform(preds) | |
| return preds | |
| def merge_four_frames(ess: pd.DataFrame, moca: pd.DataFrame, gds: pd.DataFrame, dat: pd.DataFrame) -> pd.DataFrame: | |
| ess = try_compute_ess_total(ess) | |
| gds = try_compute_gds_total(gds) | |
| for df, name in [(ess, "ESS"), (moca, "MoCA"), (gds, "GDS"), (dat, "DaTSCAN")]: | |
| if not all(col in df.columns for col in ["PATNO", "EVENT_ID"]): | |
| raise HTTPException(status_code=400, detail=f"{name} CSV missing PATNO or EVENT_ID columns") | |
| df = ess.merge(moca, on=["PATNO", "EVENT_ID"], how="inner", suffixes=("_ess", "_moca")) | |
| df = df.merge(gds, on=["PATNO", "EVENT_ID"], how="inner") | |
| df = df.merge(dat, on=["PATNO", "EVENT_ID"], how="inner") | |
| missing = [f for f in FEATURES if f not in df.columns] | |
| if missing: | |
| raise HTTPException(status_code=400, detail=f"Merged CSVs missing required features: {missing}") | |
| return df | |
| def detect_file_kind(name: str) -> str: | |
| l = name.lower() | |
| if "datscan" in l or ("dat" in l and "scan" in l): | |
| return "datscan" | |
| if "moca" in l: | |
| return "moca" | |
| if "gds" in l: | |
| return "gds" | |
| if "ess" in l: | |
| return "ess" | |
| return "unknown" | |
| # ----------------------------- | |
| # Schemas | |
| # ----------------------------- | |
| class PatientData(BaseModel): | |
| ESS_TOTAL: float | |
| MCATOT: float | |
| GDS_TOTAL: float | |
| MCAALTTM: float | |
| MCACUBE: float | |
| MCASER7: float | |
| MCAABSTR: float | |
| GDSSATIS: float | |
| GDSHAPPY: float | |
| GDSENRGY: float | |
| ESS1: float | |
| ESS2: float | |
| # ----------------------------- | |
| # Endpoints | |
| # ----------------------------- | |
| def predict_json(data: PatientData): | |
| df = pd.DataFrame([data.dict()]) | |
| preds = standardize_and_predict(df) | |
| out = preds[0].tolist() | |
| return { | |
| "predicted_biomarkers": { | |
| TARGETS[0]: out[0], | |
| TARGETS[1]: out[1], | |
| TARGETS[2]: out[2], | |
| TARGETS[3]: out[3], | |
| }, | |
| "source": "json", | |
| "rows": 1 | |
| } | |
| async def predict_files(files: List[UploadFile] = File(...)): | |
| if len(files) < 4: | |
| raise HTTPException(status_code=400, detail="Please upload four CSV files: ESS, MoCA, GDS, DaTSCAN.") | |
| buckets = {"ess": None, "moca": None, "gds": None, "datscan": None} | |
| fallback = [] | |
| for f in files: | |
| content = await f.read() | |
| try: | |
| df = pd.read_csv(io.BytesIO(content)) | |
| except Exception: | |
| raise HTTPException(status_code=400, detail=f"Could not parse CSV: {f.filename}") | |
| kind = detect_file_kind(f.filename) | |
| if kind in buckets and buckets[kind] is None: | |
| buckets[kind] = df | |
| else: | |
| fallback.append((kind, df)) | |
| if buckets["ess"] is None: | |
| candidates = [df for kind, df in fallback if "ESS1" in df.columns or "ESS_TOTAL" in df.columns] | |
| if candidates: | |
| buckets["ess"] = candidates[0] | |
| if buckets["moca"] is None: | |
| candidates = [df for kind, df in fallback if "MCATOT" in df.columns] | |
| if candidates: | |
| buckets["moca"] = candidates[0] | |
| if buckets["gds"] is None: | |
| candidates = [df for kind, df in fallback if "GDS_TOTAL" in df.columns or "GDSSATIS" in df.columns] | |
| if candidates: | |
| buckets["gds"] = candidates[0] | |
| if buckets["datscan"] is None: | |
| candidates = [df for kind, df in fallback if any(c.startswith("DATSCAN_") for c in df.columns)] | |
| if candidates: | |
| buckets["datscan"] = candidates[0] | |
| if any(v is None for v in buckets.values()): | |
| raise HTTPException(status_code=400, detail="Could not identify all four CSVs (ESS, MoCA, GDS, DaTSCAN) by filename/columns.") | |
| merged = merge_four_frames(buckets["ess"], buckets["moca"], buckets["gds"], buckets["datscan"]) | |
| preds = standardize_and_predict(merged) | |
| # Create a list of predictions for each patient | |
| results = [] | |
| for idx, pred in enumerate(preds): | |
| result = { | |
| "PATNO": int(merged.iloc[idx]["PATNO"]), | |
| "EVENT_ID": str(merged.iloc[idx]["EVENT_ID"]), | |
| "predicted_biomarkers": { | |
| TARGETS[0]: float(pred[0]), | |
| TARGETS[1]: float(pred[1]), | |
| TARGETS[2]: float(pred[2]), | |
| TARGETS[3]: float(pred[3]), | |
| } | |
| } | |
| results.append(result) | |
| return { | |
| "predictions": results, | |
| "source": "files", | |
| "merged_rows": int(merged.shape[0]) | |
| } |