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()
|