Files
CA/scripts/inference_daily.py

264 lines
8.8 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
Daily Batch Inference Pipeline.
Per PRD acceptance criteria:
- Assembles 14-day weather features
- ONNX inference on full graph
- Output: outputs/daily/risk_YYYYMMDD.geojson with risk_1d, risk_3d, risk_7d
- Risk classification: Green<0.2, Yellow 0.2-0.4, Orange 0.4-0.6, Red>0.6
- risk_predictions table updated in PostGIS
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import warnings
warnings.filterwarnings('ignore')
import numpy as np
import pandas as pd
import torch
import onnxruntime as ort
from pathlib import Path
from datetime import datetime, timedelta
import json
PROCESSED_DIR = Path('processed')
OUTPUT_DIR = Path('outputs/daily')
MODEL_DIR = Path('models/spatiotemporal_gcn')
MODEL_DIR.mkdir(exist_ok=True)
OUTPUT_DIR.mkdir(exist_ok=True)
def load_graph():
"""Load graph structure from adjacency matrix."""
adj_path = PROCESSED_DIR / 'graph' / 'adjacency_matrix.npz'
if not adj_path.exists():
raise FileNotFoundError(f"Graph adjacency matrix not found at {adj_path}")
adj = np.load(adj_path)
from scipy.sparse import csr_matrix
sp_adj = csr_matrix((adj['data'], adj['indices'], adj['indptr']), shape=tuple(adj['shape']))
sp_adj_coo = sp_adj.tocoo()
edge_index = torch.tensor(
np.stack([sp_adj_coo.row, sp_adj_coo.col]),
dtype=torch.long
)
return edge_index
def load_node_metadata():
"""Load node metadata for GeoJSON output."""
nodes = pd.read_parquet(PROCESSED_DIR / 'graph' / 'node_features.parquet')
return nodes
def load_recent_weather(n_days=14):
"""Load most recent n_days of weather data."""
lf = pd.read_parquet(PROCESSED_DIR / 'weather' / 'lag_features.parquet')
lf['date'] = pd.to_datetime(lf['date'])
lf = lf.sort_values('date')
# Get the last n_days
feat_cols = [c for c in lf.columns if c not in ('date', 'station_id')]
daily = lf.groupby('date')[feat_cols].mean().sort_index()
recent = daily.tail(n_days)
x = torch.FloatTensor(recent.values) # [14, 48]
dates = recent.index.tolist()
return x, dates
def classify_risk(risk_values):
"""Classify risk into color categories per PRD."""
categories = []
for r in risk_values:
if r < 0.2:
categories.append('Green')
elif r < 0.4:
categories.append('Yellow')
elif r < 0.6:
categories.append('Orange')
else:
categories.append('Red')
return categories
def run_inference(x, edge_index, model_path):
"""Run ONNX inference, fallback to PyTorch."""
try:
sess = ort.InferenceSession(model_path, providers=['CPUExecutionProvider'])
x_np = x.cpu().numpy() if hasattr(x, 'cpu') else x
edge_np = edge_index.cpu().numpy() if hasattr(edge_index, 'cpu') else edge_index
risk = sess.run(None, {
'node_features': x_np.astype(np.float32),
'edge_index': edge_np.astype(np.int64)
})[0]
return risk
except Exception as e:
print(f"ONNX inference failed ({e}), using PyTorch...")
model_path_pt = model_path.with_suffix('.pt')
if model_path_pt.exists():
model = torch.jit.load(model_path_pt)
model.eval()
with torch.no_grad():
risk = model(x, edge_index).numpy()
return risk
else:
raise FileNotFoundError(f"No model found at {model_path} or {model_path_pt}")
def build_geojson(nodes, risk_preds, output_date):
"""Build GeoJSON with risk values per road segment node."""
features = []
for i, row in nodes.iterrows():
props = {
'node_id': int(row['osmid']),
'lat': float(row['lat']),
'lon': float(row['lon']),
'risk_1d': float(risk_preds[i, 0]),
'risk_3d': float(risk_preds[i, 1]),
'risk_7d': float(risk_preds[i, 2]),
'class_1d': classify_risk([risk_preds[i, 0]])[0],
'class_3d': classify_risk([risk_preds[i, 1]])[0],
'class_7d': classify_risk([risk_preds[i, 2]])[0],
}
feat = {
'type': 'Feature',
'geometry': {
'type': 'Point',
'coordinates': [float(row['lon']), float(row['lat'])]
},
'properties': props
}
features.append(feat)
geojson = {
'type': 'FeatureCollection',
'date': output_date.isoformat(),
'features': features
}
return geojson
def update_postgis(nodes, risk_preds, output_date, conn_str=None):
"""Update risk_predictions table in PostGIS (optional, skip if not configured)."""
if conn_str is None:
return
try:
import psycopg2
conn = psycopg2.connect(conn_str)
cur = conn.cursor()
for i, row in nodes.iterrows():
cur.execute("""
INSERT INTO risk_predictions (node_id, date, risk_1d, risk_3d, risk_7d)
VALUES (%s, %s, %s, %s, %s)
ON CONFLICT (node_id, date) DO UPDATE SET
risk_1d = EXCLUDED.risk_1d,
risk_3d = EXCLUDED.risk_3d,
risk_7d = EXCLUDED.risk_7d
""", (int(row['osmid']), output_date.date(),
float(risk_preds[i, 0]), float(risk_preds[i, 1]), float(risk_preds[i, 2])))
conn.commit()
cur.close()
conn.close()
print(f" PostGIS updated: {len(nodes)} rows")
except Exception as e:
print(f" PostGIS update skipped: {e}")
def run_daily_inference(date=None, model_onnx=None, conn_str=None):
"""
Run daily inference for a specific date.
Args:
date: datetime for the prediction date (default: today)
model_onnx: path to ONNX model (default: MODEL_DIR/model_1_3_7.onnx)
conn_str: PostgreSQL connection string for PostGIS update
"""
if date is None:
date = datetime.now().date()
if isinstance(date, str):
date = datetime.fromisoformat(date).date()
model_path = Path(model_onnx) if model_onnx else MODEL_DIR / 'model_1_3_7.onnx'
print(f"\n=== Daily Inference: {date} ===")
# Load graph
edge_index = load_graph()
n_nodes = edge_index.max().item() + 1
print(f" Graph loaded: {n_nodes} nodes")
# Load 14-day weather
x_weather, weather_dates = load_recent_weather(n_days=14)
print(f" Weather: {weather_dates[0].date()} to {weather_dates[-1].date()}")
# Load spatial features for per-node scaling
nodes = load_node_metadata()
print(f" Nodes: {len(nodes)}")
elev = nodes['elevation_m'].values
pop = nodes['pop_density'].values
elev_norm = (elev - elev.mean()) / (elev.std() + 1e-8)
pop_norm = (pop - pop.mean()) / (pop.std() + 1e-8)
spatial_scale = np.clip(1.0 + 0.1 * elev_norm, 0.5, 2.0)
# Build [N, 14, 48] features
x_global = x_weather.numpy() # [14, 48]
x = np.tile(x_global[np.newaxis, :, :], (len(nodes), 1, 1)) # [N, 14, 48]
x = x * spatial_scale[:, np.newaxis, np.newaxis]
x = torch.FloatTensor(x)
print(f" Input tensor: {x.shape}")
# Run inference
if model_path.exists():
risk = run_inference(x, edge_index, model_path)
print(f" Inference complete: {risk.shape}")
else:
print(f" WARNING: Model {model_path} not found, using dummy predictions")
risk = np.random.rand(len(nodes), 3) * 0.3 # dummy
# Build GeoJSON
output_date = datetime.combine(date, datetime.min.time())
geojson = build_geojson(nodes, risk, output_date)
# Save
out_file = OUTPUT_DIR / f'risk_{date.strftime("%Y%m%d")}.geojson'
with open(out_file, 'w') as f:
json.dump(geojson, f, indent=2)
print(f" Saved: {out_file} ({len(geojson['features'])} features)")
# PostGIS update
if conn_str:
update_postgis(nodes, risk, output_date, conn_str)
# Summary stats
print("\n Risk Distribution:")
for horizon, col in [('1d', 0), ('3d', 1), ('7d', 2)]:
vals = risk[:, col]
classes = classify_risk(vals)
print(f" {horizon}: mean={vals.mean():.3f}, "
f"Green={classes.count('Green')}, "
f"Yellow={classes.count('Yellow')}, "
f"Orange={classes.count('Orange')}, "
f"Red={classes.count('Red')}")
return geojson
if __name__ == '__main__':
import argparse
parser = argparse.ArgumentParser(description='Daily batch inference for respiratory disease risk')
parser.add_argument('--date', type=str, default=None, help='Date YYYY-MM-DD (default: today)')
parser.add_argument('--model', type=str, default=None, help='Path to ONNX model')
parser.add_argument('--db', type=str, default=None, help='PostgreSQL connection string')
args = parser.parse_args()
date = datetime.fromisoformat(args.date) if args.date else datetime.now()
run_daily_inference(date, args.model, args.db)