Files
CA/scripts/generate_grid_features.py

130 lines
4.7 KiB
Python
Raw Permalink Normal View History

#!/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()