Context: Build a spatial risk assessment system correlating air quality data with children's respiratory disease incidence across Wuhan. Approach: FastAPI backend serving PostGIS spatial queries, React frontend with Deck.gl maps, and a PyTorch SpatialTemporalGCN pipeline for multi-day (1d/3d/7d) risk prediction. Changes: - backend/ — FastAPI API with auth (JWT), alerts, risk analysis, geocoded case data, grid statistics, and report endpoints - frontend/ — React dashboard with interactive risk maps, alert monitoring, district comparison charts, and timeline player - models/ — SpatialTemporalGCN model with trained weights and ONNX export for inference - scripts/ — ETL pipeline for weather + medical data, grid generation, feature engineering, training, and daily inference - deploy/ — Docker Compose configs for backend, frontend, and MLflow - docs/ — API docs, deployment guide, user guide, and code review Impact: Enables spatial risk visualization, alert monitoring, and ML-driven health risk forecasting for environmental health teams.
130 lines
4.7 KiB
Python
130 lines
4.7 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
Grid Feature Generator for ML Model
|
|
Generates features on-demand for model inference.
|
|
|
|
Strategy:
|
|
- Weather: Interpolate from stations to grid on-demand
|
|
- Cases: Use district-level aggregation (already computed)
|
|
- DEM/Pop: Static features from resampled rasters
|
|
|
|
Usage:
|
|
python scripts/generate_grid_features.py --date 2022-01-01 --output processed/features_2022-01-01.parquet
|
|
"""
|
|
|
|
import pandas as pd
|
|
import numpy as np
|
|
from scipy.interpolate import griddata
|
|
from pathlib import Path
|
|
import argparse
|
|
import time
|
|
|
|
class GridFeatureGenerator:
|
|
def __init__(self):
|
|
print("Loading static data...")
|
|
|
|
# Load 100m grid index (998,601 cells)
|
|
# Support running from backend/ directory
|
|
self.base_path = Path(__file__).parent.parent
|
|
self.grid_df = pd.read_parquet(self.base_path / 'processed/grid_100m_index.parquet')
|
|
self.grid_points = self.grid_df[['center_lon', 'center_lat']].values
|
|
self.grid_ids = self.grid_df['grid_id'].values
|
|
print(f" Grid: {len(self.grid_ids):,} cells")
|
|
|
|
# Load district mapping
|
|
self.district_map = pd.read_parquet(self.base_path / 'processed/grid_district_mapping.parquet')
|
|
print(f" District mapping: {len(self.district_map):,} rows")
|
|
|
|
# Station data cache
|
|
self.station_cache = {}
|
|
|
|
def load_station_data(self, date_str):
|
|
date = pd.to_datetime(date_str).date()
|
|
year = date.year
|
|
|
|
if year not in self.station_cache:
|
|
self.station_cache[year] = pd.read_parquet(f'processed/weather/station_daily_{year}.parquet')
|
|
self.station_cache[year]['date'] = pd.to_datetime(self.station_cache[year]['date']).dt.date
|
|
|
|
station_df = self.station_cache[year]
|
|
day_data = station_df[station_df['date'] == date]
|
|
|
|
if len(day_data) == 0:
|
|
raise ValueError(f"No station data for {date}")
|
|
|
|
return day_data
|
|
|
|
def interpolate_weather(self, day_data, pollutant):
|
|
stations = day_data[['lon', 'lat', pollutant]].dropna()
|
|
|
|
if len(stations) < 3:
|
|
return np.full(len(self.grid_ids), np.nan)
|
|
|
|
result = griddata(
|
|
stations[['lon', 'lat']].values,
|
|
stations[pollutant].values,
|
|
self.grid_points,
|
|
method='nearest'
|
|
)
|
|
|
|
return result
|
|
|
|
def get_cases_for_date(self, date_str):
|
|
date = pd.to_datetime(date_str).date()
|
|
cases_df = pd.read_parquet(self.base_path / 'processed/cases_by_district_daily.parquet')
|
|
cases_df['date'] = pd.to_datetime(cases_df['date']).dt.date
|
|
|
|
day_cases = cases_df[cases_df['date'] == date]
|
|
merged = self.district_map.merge(day_cases, left_on='district_name', right_on='district', how='left')
|
|
|
|
return merged
|
|
|
|
def generate_features(self, date_str):
|
|
print(f"Generating features for {date_str}...")
|
|
t0 = time.time()
|
|
|
|
# Load weather data
|
|
day_data = self.load_station_data(date_str)
|
|
|
|
# Interpolate pollutants to grid
|
|
pollutants = ['AQI', 'PM25', 'PM10', 'SO2', 'NO2', 'O3', 'CO']
|
|
features = {'grid_id': self.grid_ids}
|
|
|
|
for poll in pollutants:
|
|
print(f" Interpolating {poll}...")
|
|
features[poll] = self.interpolate_weather(day_data, poll)
|
|
|
|
# Add case data by district
|
|
print(" Adding case data...")
|
|
cases_merged = self.get_cases_for_date(date_str)
|
|
features['outpatient_count'] = cases_merged['outpatient_count'].fillna(0).values
|
|
features['inpatient_count'] = cases_merged['inpatient_count'].fillna(0).values
|
|
features['total_cases'] = cases_merged['total_cases'].fillna(0).values
|
|
features['district'] = cases_merged['district_name'].values
|
|
|
|
feature_df = pd.DataFrame(features)
|
|
feature_df['date'] = date_str
|
|
|
|
print(f"Generated {len(feature_df):,} rows in {time.time()-t0:.1f}s")
|
|
return feature_df
|
|
|
|
def save_features(self, feature_df, output_path):
|
|
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
|
|
feature_df.to_parquet(output_path, index=False, compression='gzip')
|
|
print(f"Saved: {output_path}")
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description='Generate grid features for ML model')
|
|
parser.add_argument('--date', required=True, help='Date (YYYY-MM-DD)')
|
|
parser.add_argument('--output', required=True, help='Output parquet path')
|
|
args = parser.parse_args()
|
|
|
|
generator = GridFeatureGenerator()
|
|
features = generator.generate_features(args.date)
|
|
generator.save_features(features, args.output)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
main()
|