IDKHowToCodeFr/tinyml-backend
1
1import sys2import os3sys.path.append(os.path.dirname(os.path.abspath(__file__)))4 5from fastapi import FastAPI, UploadFile, File, BackgroundTasks, HTTPException, WebSocket, WebSocketDisconnect6from fastapi.middleware.cors import CORSMiddleware7import pandas as pd8import asyncio9import shap10import database as db11import logging12from contextlib import asynccontextmanager13import warnings14 15# Suppress sklearn version mismatch warnings in stdout16try:17 from sklearn.exceptions import InconsistentVersionWarning18 warnings.filterwarnings("ignore", category=InconsistentVersionWarning)19except ImportError:20 pass21 22from preprocessing import preprocess_data23from export import generate_c_code24from inference import evaluate25from schemas import PatientData26from mlops import MLOpsEngine27from streamer import TelemetryStreamer28from fastapi.responses import RedirectResponse29 30# Setup standard logging for critical fault alerts31logging.basicConfig(filename="alerts.log", level=logging.WARNING, 32 format='%(asctime)s %(levelname)s: %(message)s', datefmt='%Y-%m-%d %H:%M:%S')33 34ml_engine = None35 36async def background_sync():37 while True:38 await asyncio.sleep(60)39 await db.sync_from_hub()40 await db.sync_to_hub()41 42@asynccontextmanager43async def lifespan(app: FastAPI):44 global ml_engine45 await db.init_db()46 sync_task = asyncio.create_task(background_sync())47 48 data_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), '..', 'data', 'patient_dataset.csv')49 ml_engine = MLOpsEngine(data_path)50 51 yield52 sync_task.cancel()53 ml_engine = None54 55app = FastAPI(title="TinyML Healthcare API", lifespan=lifespan)56 57app.add_middleware(58 CORSMiddleware,59 allow_origins=["*"],60 allow_credentials=True,61 allow_methods=["*"],62 allow_headers=["*"],63)64 65@app.get("/")66def read_root():67 return RedirectResponse(url="/docs")68 69@app.get("/health")70def health_check():71 return {"status": "Healthy" if ml_engine and ml_engine.get_ensemble() else "Warning - Models Offline"}72 73async def log_prediction_background(data, final_pred, conf):74 try:75 await db.log_prediction(data, final_pred, conf)76 except Exception as db_e:77 print(f"DB logging skipped: {db_e}")78 79@app.post("/predict")80async def predict(data: PatientData, background_tasks: BackgroundTasks):81 try:82 eng = ml_engine.get_ensemble()83 if not eng:84 return {"error": "Models untrained. Ensure python backend/models.py executes."}85 86 result = await asyncio.to_thread(evaluate, eng, data)87 if "error" in result:88 return result89 90 is_at_risk = result["prediction"]91 final_pred = result["prediction_label"]92 conf = result["probability"]93 94 # Critical Alert System — standard library logging95 if is_at_risk == 1 and float(conf) > 0.80:96 logging.warning(f"Patient at risk! HR: {data.Heart_Rate}, SpO2: {data.SpO2_Level}, Confidence: {conf:.2f}")97 98 # Log to SQLite History99 background_tasks.add_task(log_prediction_background, data, final_pred, float(conf))100 101 return result102 except Exception as e:103 import traceback104 raise HTTPException(status_code=500, detail=f"Backend Error: {str(e)}\n\nTraceback:\n{traceback.format_exc()}")105 106@app.websocket("/ws/feed")107async def websocket_feed(websocket: WebSocket):108 await websocket.accept()109 eng = ml_engine.get_ensemble()110 if not eng:111 await websocket.close(code=1011)112 return113 114 streamer = TelemetryStreamer(eng)115 try:116 async for payload in streamer.stream():117 await websocket.send_json(payload)118 except WebSocketDisconnect:119 pass120 except Exception as e:121 print(f"WS error: {e}")122 123@app.get("/history")124async def history():125 return await db.get_history()126 127@app.get("/dataset")128def get_dataset():129 if ml_engine and os.path.exists(ml_engine.dataset_path):130 try:131 df = pd.read_csv(ml_engine.dataset_path, encoding='utf-8')132 df.columns = [c.strip() for c in df.columns]133 return df.to_dict(orient="records")134 except Exception as e:135 return {"error": f"Failed to read dataset: {str(e)}"}136 return {"error": "Dataset not found"}137 138@app.get("/sync")139def force_sync():140 from database import sync_from_hub, sync_to_hub141 sync_from_hub()142 return {"status": "Sync attempted"}143 144@app.post("/explain")145async def explain(data: PatientData):146 eng = ml_engine.get_ensemble() if ml_engine else None147 if not eng or 'rf' not in eng.models:148 return {"error": "RF Model unavailable for explanation."}149 150 df = pd.DataFrame([data.model_dump()])151 X_proc, _ = await asyncio.to_thread(preprocess_data, df, False)152 153 def compute_shap():154 import numpy as np155 rf_model = eng.models['rf']156 explainer = shap.TreeExplainer(rf_model)157 shap_values = explainer.shap_values(X_proc)158 pred_idx = int(rf_model.predict(X_proc)[0])159 if isinstance(shap_values, list):160 vals = shap_values[pred_idx][0]161 elif isinstance(shap_values, np.ndarray) and len(shap_values.shape) == 3:162 vals = shap_values[0, :, pred_idx]163 else:164 vals = shap_values[0]165 return vals.tolist(), X_proc.columns.tolist()166 167 try:168 shap_vals, features = await asyncio.to_thread(compute_shap)169 return {"shap_values": shap_vals, "feature_names": features}170 except Exception as e:171 return {"error": str(e)}172 173@app.get("/export_tinyml")174def export_tinyml(model_name: str = "rf", quantize: bool = False):175 eng = ml_engine.get_ensemble() if ml_engine else None176 if not eng:177 return {"error": "Models untrained or engine offline"}178 return generate_c_code(eng, model_name, quantize)179 180def _run_mlops_ingest(csv_bytes: bytes):181 return ml_engine.ingest_batch_sync(csv_bytes)182 183@app.post("/retrain")184async def retrain(background_tasks: BackgroundTasks, file: UploadFile = File(...)):185 content = await file.read()186 187 # We add this to background tasks so the HTTP response returns immediately188 # while the engine merges CSVs and retrains the models.189 def bg_task():190 res = ml_engine.ingest_batch_sync(content)191 if not res.success:192 print(f"MLOps Background Task Failed: {res.message}")193 194 background_tasks.add_task(bg_task)195 return {"status": "success", "message": "Dataset uploaded and ensemble retrain started in background!"}196 