CoolFace
Apppublic

Vikctor/Drought_Disaster_Models

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py253 linesDownload Raw Back to root
1from fastapi import FastAPI, HTTPException, Request, Response2from fastapi.middleware.cors import CORSMiddleware3from fastapi.openapi.docs import get_swagger_ui_html, get_redoc_html4from pydantic import BaseModel5import pandas as pd6import joblib7import requests8import gc9import os10import logging11from math import sin, cos, radians, pi12from contextlib import asynccontextmanager13 14# -------------------------15# Logger16# -------------------------17logging.basicConfig(18    level=logging.INFO,19    format="%(asctime)s - %(levelname)s - %(message)s"20)21 22# -------------------------23# Global models24# -------------------------25_occurrence_model = None26_occurrence_scaler = None27_severity_model = None28_severity_scaler = None29 30# -------------------------31# Feature setup32# -------------------------33API_BASE = "https://power.larc.nasa.gov/api/temporal/daily/point"34PARAMS = "PRECTOT,T2M,T2M_MAX,T2M_MIN,ALLSKY_SFC_SW_DWN,RH2M,WS2M"35FEATURE_ORDER = [36    "RH2M", "T2M_MAX", "T2M_MIN", "WS2M", "T2M",37    "ALLSKY_SFC_SW_DWN", "PRECTOTCORR",38    "lat_sin", "lat_cos", "lon_sin", "lon_cos",39    "month_sin", "month_cos"40]41 42# -------------------------43# Utility functions44# -------------------------45def cleanup_memory():46    gc.collect()47 48def safe_model_load(filename: str):49    try:50        script_dir = os.path.dirname(os.path.abspath(__file__))51        path = os.path.join(script_dir, filename)52        if not os.path.exists(path):53            raise FileNotFoundError(f"{filename} not found")54        return joblib.load(path)55    except Exception as e:56        logging.error(f"Failed to load {filename}: {e}")57        raise HTTPException(status_code=500, detail=f"Model loading failed: {filename}")58 59def get_occurrence_model_and_scaler():60    global _occurrence_model, _occurrence_scaler61    if _occurrence_model is None or _occurrence_scaler is None:62        logging.info("Loading occurrence model/scaler...")63        _occurrence_model = safe_model_load("drought_occurrence_model.joblib")64        _occurrence_scaler = safe_model_load("drought_occurrence_model_scaler.joblib")65        cleanup_memory()66    return _occurrence_model, _occurrence_scaler67 68def get_severity_model_and_scaler():69    global _severity_model, _severity_scaler70    if _severity_model is None or _severity_scaler is None:71        logging.info("Loading severity model/scaler...")72        _severity_model = safe_model_load("drought_severity_model.joblib")73        _severity_scaler = safe_model_load("drought_severity_model_scaler.joblib")74        cleanup_memory()75    return _severity_model, _severity_scaler76 77# -------------------------78# Lifespan79# -------------------------80@asynccontextmanager81async def lifespan(app: FastAPI):82    logging.info("๐Ÿš€ Drought API starting (models load on first request)")83    cleanup_memory()84    yield85    logging.info("๐Ÿ›‘ Shutting down API")86    global _occurrence_model, _occurrence_scaler, _severity_model, _severity_scaler87    _occurrence_model = _occurrence_scaler = _severity_model = _severity_scaler = None88    cleanup_memory()89 90# -------------------------91# FastAPI instance92# -------------------------93app = FastAPI(94    title="๐ŸŒ Drought Prediction API",95    version="2.4",96    description="Memory-optimized drought prediction API",97    lifespan=lifespan98)99 100# -------------------------101# CORS middleware for website102# -------------------------103app.add_middleware(104    CORSMiddleware,105    allow_origins=["*"],  # replace with website URL in production106    allow_methods=["*"],107    allow_headers=["*"]108)109 110# -------------------------111# Request model112# -------------------------113class PredictionRequest(BaseModel):114    lat: float115    lon: float116    time: str  # YYYY-MM-DD117 118# -------------------------119# NASA feature fetcher120# -------------------------121def fetch_features(lat, lon, time_str: str) -> dict:122    end = pd.to_datetime(time_str)123    start = end - pd.Timedelta(days=90)124    params = {125        "latitude": lat,126        "longitude": lon,127        "start": start.strftime("%Y%m%d"),128        "end": end.strftime("%Y%m%d"),129        "parameters": PARAMS,130        "format": "JSON",131        "community": "AG"132    }133    try:134        response = requests.get(API_BASE, params=params, timeout=30)135        response.raise_for_status()136        data = response.json().get("properties", {}).get("parameter", {})137        features = {}138        for p, vals in data.items():139            values = [v for v in vals.values() if v is not None]140            if values:141                features["PRECTOTCORR" if p=="PRECTOT" else p] = sum(values)/len(values) if p!="PRECTOT" else sum(values)142        features.update({143            "lat_sin": sin(radians(lat)),144            "lat_cos": cos(radians(lat)),145            "lon_sin": sin(radians(lon)),146            "lon_cos": cos(radians(lon)),147            "month_sin": sin(2*pi*end.month/12),148            "month_cos": cos(2*pi*end.month/12)149        })150        missing = [f for f in FEATURE_ORDER if f not in features]151        if missing:152            raise HTTPException(status_code=500, detail=f"Missing features: {missing}")153        cleanup_memory()154        return features155    except Exception as e:156        logging.error(f"NASA fetch error: {e}")157        raise HTTPException(status_code=502, detail="NASA API request failed")158 159# -------------------------160# Prediction endpoint161# -------------------------162@app.post("/predict")163async def predict(req: PredictionRequest):164    try:165        features = fetch_features(req.lat, req.lon, req.time)166        X = pd.DataFrame([[features[f] for f in FEATURE_ORDER]], columns=FEATURE_ORDER)167        occ_model, occ_scaler = get_occurrence_model_and_scaler()168        sev_model, sev_scaler = get_severity_model_and_scaler()169        X_occ = occ_scaler.transform(X)170        X_sev = sev_scaler.transform(X)171        occurrence_pred = int(occ_model.predict(X_occ)[0])172        occurrence_proba = occ_model.predict_proba(X_occ)[0].tolist()173        severity_pred = int(sev_model.predict(X_sev)[0])174        severity_proba = sev_model.predict_proba(X_sev)[0].tolist()175        result = {176            "input": {"lat": req.lat, "lon": req.lon, "time": req.time},177            "occurrence": {"prediction": occurrence_pred, "probabilities": occurrence_proba},178            "severity": {"prediction": severity_pred, "probabilities": severity_proba},179            "features_used": {k: round(v,4) for k,v in zip(FEATURE_ORDER, X.iloc[0].tolist())}180        }181        cleanup_memory()182        return result183    except HTTPException as e:184        raise e185    except Exception as e:186        logging.error(f"Prediction error: {e}")187        raise HTTPException(status_code=500, detail=str(e))188 189# -------------------------190# Health check191# -------------------------192@app.api_route("/health", methods=["GET", "HEAD"])193async def health_check(request: Request):194    if request.method == "HEAD":195        return Response(status_code=200)196    return {"status": "healthy", "api_version": "2.4"}197 198# -------------------------199# Debug endpoint200# -------------------------201@app.get("/debug")202async def debug_info():203    return {204        "models_loaded": {205            "occurrence_model": _occurrence_model is not None,206            "occurrence_scaler": _occurrence_scaler is not None,207            "severity_model": _severity_model is not None,208            "severity_scaler": _severity_scaler is not None209        },210        "feature_order": FEATURE_ORDER211    }212 213# -------------------------214# Test endpoint215# -------------------------216@app.get("/test")217async def test_prediction():218    try:219        test_req = PredictionRequest(lat=40.7128, lon=-74.0060, time="2024-08-15")220        result = await predict(test_req)221        return {"test_status": "success", "result": result}222    except Exception as e:223        return {"test_status": "failed", "error": str(e)}224 225# -------------------------226# Root endpoint227# -------------------------228@app.get("/")229async def root():230    return {231        "message": "๐ŸŒ Drought Prediction API",232        "version": "2.4",233        "endpoints": {234            "predict": "/predict",235            "health": "/health",236            "debug": "/debug",237            "test": "/test",238            "docs": "/docs",239            "redoc": "/redoc"240        }241    }242 243# -------------------------244# Swagger UI and Redoc245# -------------------------246@app.get("/docs", include_in_schema=False)247async def custom_swagger_ui():248    return get_swagger_ui_html(openapi_url="/openapi.json", title="API Docs")249 250@app.get("/redoc", include_in_schema=False)251async def custom_redoc():252    return get_redoc_html(openapi_url="/openapi.json", title="ReDoc")253