Team Ai
Apppublic

IDKHowToCodeFr/tinyml-backend

sourceHugging Faceupdated 7h agoView on Hugging Face
1likes
main.py196 linesDownload Raw Back to backend
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