Vikctor/Drought_Disaster_Models
0
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 