Files
CA/scripts/train_model.py
Akiba So fc468464b2 feat: Initial CBPOA commit — 武汉儿童呼吸疾病风险评估系统
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.
2026-06-05 02:13:49 +08:00

461 lines
16 KiB
Python

#!/usr/bin/env python3
"""
Training Pipeline for Spatial-Temporal Transformer + GCN Model.
Simplified approach: use global weather mean per day as node features,
scaled by per-node spatial features (elevation, population density).
"""
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 torch.nn as nn
import torch.optim as optim
import mlflow
from pathlib import Path
from datetime import datetime
from models.spatiotemporal_gcn.model import SpatialTemporalGCN
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {DEVICE}")
PROCESSED_DIR = Path('processed')
MODEL_DIR = Path('models/spatiotemporal_gcn')
MODEL_DIR.mkdir(exist_ok=True)
LEARNING_RATE = 1e-4
WEIGHT_DECAY = 0.01
PATIENCE = 15
MAX_EPOCHS = 200
BATCH_SIZE = 1024
# Data split (medical data only available in December)
TRAIN_START = '2022-12-01'
TRAIN_END = '2022-12-31'
VAL_START = '2023-12-01'
VAL_END = '2023-12-31'
# Baseline MAE from compute_baseline_mae.py
BASELINE_MAE = {'1-day': 0.2314, '3-day': 0.5424, '7-day': 0.6391}
SEED = 42
np.random.seed(SEED)
torch.manual_seed(SEED)
def load_all_data():
"""Load all processed data."""
print("Loading data...")
# Graph
adj = np.load(PROCESSED_DIR / 'graph' / 'adjacency_matrix.npz')
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
) # Keep on CPU for subgraph operations
# Node metadata
nodes = pd.read_parquet(PROCESSED_DIR / 'graph' / 'node_features.parquet')
n_nodes = len(nodes)
print(f" Graph: {n_nodes} nodes, {edge_index.shape[1]} edges")
# Weather lag features (station-level daily)
lf = pd.read_parquet(PROCESSED_DIR / 'weather' / 'lag_features.parquet')
lf['date'] = pd.to_datetime(lf['date'])
lf = lf.sort_values('date')
print(f" Weather: {len(lf)} records, {lf['station_id'].nunique()} stations")
# Medical targets (district-level daily)
out = pd.read_csv(PROCESSED_DIR / 'medical' / 'outpatient_daily.csv', parse_dates=['date'])
inp = pd.read_csv(PROCESSED_DIR / 'medical' / 'inpatient_daily.csv', parse_dates=['date'])
out['weight'] = 1
inp['weight'] = 3
combined = pd.concat([out, inp])
combined['weighted_cases'] = combined['case_count'] * combined['weight']
medical = combined.groupby(['date', 'district']).agg(
weighted_cases=('weighted_cases', 'sum')
).reset_index()
medical['risk'] = medical.groupby('district')['weighted_cases'].transform(
lambda x: x / x.mean()
)
print(f" Medical: {len(medical)} district-day records")
return edge_index, nodes, lf, medical
def build_global_weather_timeseries(lf):
"""
Build global mean weather per day: [T, 48]
"""
feat_cols = [c for c in lf.columns if c not in ('date', 'station_id')]
daily_mean = lf.groupby('date')[feat_cols].mean()
daily_mean = daily_mean.sort_index()
dates = daily_mean.index.tolist()
x_global = daily_mean.values.astype(np.float32) # [T, 48]
return x_global, dates
def build_node_targets(nodes, medical, dates):
"""
Build per-node risk target per day: [N, T]
Use district-level medical risk, tiled to all nodes in district.
District assignment from node lat/lon nearest centroid (simplified: use 'unknown').
For nodes with no district match, use global mean risk.
"""
n_nodes = len(nodes)
n_days = len(dates)
# Global mean risk per day
global_risk = medical.groupby('date')['risk'].mean()
global_risk_dict = global_risk.to_dict()
# For each node, assign a district based on nearest centroid
# (simplified: just use global risk for all nodes)
targets = np.full((n_nodes, n_days), np.nan, dtype=np.float32)
for i, d in enumerate(dates):
if d in global_risk_dict:
targets[:, i] = global_risk_dict[d]
# Normalize per node
node_means = np.nanmean(targets, axis=1, keepdims=True)
node_means[node_means == 0] = 1
targets = targets / (node_means + 1e-8)
return targets, dates
def build_spatial_scalars(nodes):
"""
Pre-compute per-node spatial scaling factors (small, O(N)).
Returns: elev_scale [N], pop_scale [N]
"""
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)
# Scaling factors
elev_scale = 1.0 + 0.1 * elev_norm
elev_scale = np.clip(elev_scale, 0.5, 2.0).astype(np.float32)
# Pop scale (optional, can be 1.0 if not used)
pop_scale = np.ones_like(elev_scale) # or add similar modulation if needed
return elev_scale, pop_scale
def get_batch_features(elev_scale, x_global, node_indices):
"""
Compute features for a batch of nodes on-the-fly.
elev_scale: [N] pre-computed spatial scalars
x_global: [T, F] global weather per day
node_indices: list/array of node indices to fetch
Returns: [len(node_indices), T, F]
"""
batch_size = len(node_indices)
T, F = x_global.shape
# Get spatial scales for batch
batch_elev = elev_scale[node_indices]
# Tile global weather for batch: [T, F] -> [batch, T, F]
x = np.tile(x_global[np.newaxis, :, :], (batch_size, 1, 1))
# Apply spatial scaling
x = x * batch_elev[:, np.newaxis, np.newaxis]
return x.astype(np.float32)
def make_time_windows_lazy(x_global, elev_scale, targets, dates, window=14):
"""
Create time window metadata without materializing full [N, T, F] tensor.
Returns list of (time_start, node_indices) tuples for lazy feature fetching.
x_global: [T, F] global weather
elev_scale: [N] spatial scaling per node
targets: [N, T] target values
window: input window size
"""
N = len(elev_scale)
T = x_global.shape[0]
# Store window metadata: which time steps and which nodes
windows_meta = []
for t in range(T - window + 1):
# All nodes for this time window
windows_meta.append({
'time_start': t,
'time_end': t + window,
'target_time': t + window - 1,
})
return windows_meta
def train_epoch_lazy(model, windows_meta, x_global, elev_scale, y, edge_index,
optimizer, criterion, batch_size=1024):
"""
Train one epoch using lazy feature computation.
For each window, sample a batch of nodes and compute features on-the-fly.
"""
from torch_geometric.utils import subgraph
model.train()
total_loss = 0
n_batches = 0
n_windows = len(windows_meta)
n_nodes = len(elev_scale)
# Process each time window
for window_meta in windows_meta:
t_start = window_meta['time_start']
t_end = window_meta['time_end']
t_target = window_meta['target_time']
node_indices = np.random.choice(n_nodes, size=min(batch_size, n_nodes), replace=False)
node_indices_torch = torch.tensor(node_indices, dtype=torch.long)
x_batch = get_batch_features(elev_scale, x_global[t_start:t_end], node_indices)
y_batch = y[node_indices, t_target]
# Filter out NaN targets
valid_mask = ~np.isnan(y_batch)
if valid_mask.sum() == 0:
continue
# Extract subgraph and manually remap indices to ensure correctness
sub_edge_index, edge_mask = subgraph(node_indices_torch, edge_index, relabel_nodes=False)
# Create remapping: global_id -> local_idx (0 to batch_size-1)
# Use index_put for efficient remapping
local_idx = torch.arange(len(node_indices), dtype=torch.long)
remap_tensor = torch.full((n_nodes,), -1, dtype=torch.long)
remap_tensor[node_indices_torch] = local_idx
# Remap edge indices
sub_edge_index = remap_tensor[sub_edge_index]
# Validate: all indices should be in [0, batch_size)
assert sub_edge_index.min() >= 0 and sub_edge_index.max() < len(node_indices), \
f"Edge index out of bounds: min={sub_edge_index.min()}, max={sub_edge_index.max()}"
x_batch = torch.FloatTensor(x_batch).to(DEVICE)
y_batch = torch.FloatTensor(y_batch).to(DEVICE)
sub_edge_index = sub_edge_index.to(DEVICE)
# Forward pass - only compute loss on valid samples
optimizer.zero_grad()
out = model(x_batch, sub_edge_index)
loss = criterion(out[valid_mask, 1], y_batch[valid_mask])
loss.backward()
optimizer.step()
total_loss += loss.item()
n_batches += 1
return total_loss / max(n_batches, 1)
def evaluate_lazy(model, x_global, elev_scale, y, edge_index, window=14):
"""
Evaluate MAE per horizon using lazy feature computation.
"""
from torch_geometric.utils import subgraph
model.eval()
T = x_global.shape[0]
n_nodes = len(elev_scale)
horizons = {'1-day': 1, '3-day': 3, '7-day': 7}
results = {}
with torch.no_grad():
for name, h in horizons.items():
if h > T - window:
results[name] = float('nan')
continue
preds_all = []
acts_all = []
for t in range(window, T - h + 1, 5):
node_indices = np.random.choice(n_nodes, size=min(100, n_nodes), replace=False)
node_indices_torch = torch.tensor(node_indices, dtype=torch.long)
x_win = get_batch_features(elev_scale, x_global[t-window:t], node_indices)
x_win = torch.FloatTensor(x_win).to(DEVICE)
y_actual = y[node_indices, t+h-1]
y_actual = torch.FloatTensor(y_actual).to(DEVICE)
# Extract subgraph and manually remap indices
sub_edge_index, _ = subgraph(node_indices_torch, edge_index, relabel_nodes=False)
# Remap global IDs to local indices
local_idx = torch.arange(len(node_indices), dtype=torch.long)
remap_tensor = torch.full((n_nodes,), -1, dtype=torch.long)
remap_tensor[node_indices_torch] = local_idx
sub_edge_index = remap_tensor[sub_edge_index]
sub_edge_index = sub_edge_index.to(DEVICE)
# Filter out NaN targets
valid_mask = ~torch.isnan(y_actual)
if valid_mask.sum() == 0:
continue
pred = model(x_win, sub_edge_index)[valid_mask, 1]
preds_all.append(pred.mean())
acts_all.append(y_actual[valid_mask].mean())
if preds_all:
preds = torch.stack(preds_all).mean()
acts = torch.stack(acts_all).mean()
results[name] = torch.mean(torch.abs(preds - acts)).item()
else:
results[name] = float('nan')
return results
def main():
print(f"\n=== Training Pipeline === {datetime.now()}")
edge_index, nodes, lf, medical = load_all_data()
n_nodes = len(nodes)
# Build time series - global weather only (small: [T, 48])
x_global, weather_dates = build_global_weather_timeseries(lf)
targets, _ = build_node_targets(nodes, medical, weather_dates)
# Pre-compute spatial scalars (small: O(N))
elev_scale, pop_scale = build_spatial_scalars(nodes)
print(f"\nGlobal weather: {x_global.shape} [T, F]")
print(f"Targets: {targets.shape} [N, T]")
print(f"Spatial scalars: {len(elev_scale)} nodes")
# Align to training period
dates_arr = pd.to_datetime(weather_dates)
train_mask = (dates_arr >= TRAIN_START) & (dates_arr <= TRAIN_END)
val_mask = (dates_arr >= VAL_START) & (dates_arr <= VAL_END)
x_global_train = x_global[train_mask]
y_train = targets[:, train_mask]
x_global_val = x_global[val_mask]
y_val = targets[:, val_mask]
train_days = len(x_global_train)
val_days = len(x_global_val)
print(f"Train: {train_days} steps, Val: {val_days} steps")
# Make training windows (metadata only, no large tensors)
WINDOW = 14
print("Building training windows (metadata)...")
windows_meta = make_time_windows_lazy(x_global_train, elev_scale, y_train,
pd.to_datetime(weather_dates)[train_mask].tolist(),
window=WINDOW)
print(f" {len(windows_meta)} windows")
# Model
model = SpatialTemporalGCN(
node_features=48,
temporal_heads=4,
temporal_layers=3,
gcn_hidden=128,
gcn_output=64,
dropout=0.2
).to(DEVICE)
print(f"\nModel params: {sum(p.numel() for p in model.parameters()):,}")
optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
criterion = nn.L1Loss()
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=5, factor=0.5)
# Truncate edge_index for batch processing
edge_idx_trunc = edge_index[:, :min(edge_index.shape[1], n_nodes * 4)].contiguous()
mlflow.set_experiment("wuhan_respiratory_training")
with mlflow.start_run(run_name=f"train_{datetime.now().strftime('%Y%m%d_%H%M%S')}"):
mlflow.log_params({
"learning_rate": LEARNING_RATE,
"weight_decay": WEIGHT_DECAY,
"patience": PATIENCE,
"max_epochs": MAX_EPOCHS,
"window": WINDOW,
"n_nodes": n_nodes,
"train_start": TRAIN_START, "train_end": TRAIN_END,
"val_start": VAL_START, "val_end": VAL_END,
"baseline_mae_1d": BASELINE_MAE['1-day'],
"baseline_mae_3d": BASELINE_MAE['3-day'],
"baseline_mae_7d": BASELINE_MAE['7-day'],
})
best_val_mae = float('inf')
patience_counter = 0
best_state = None
for epoch in range(1, MAX_EPOCHS + 1):
train_loss = train_epoch_lazy(model, windows_meta, x_global_train, elev_scale,
y_train, edge_idx_trunc, optimizer, criterion, BATCH_SIZE)
val_mae_h = evaluate_lazy(model, x_global_val, elev_scale, y_val, edge_idx_trunc, WINDOW)
val_mae = np.nanmean(list(val_mae_h.values()))
scheduler.step(val_mae)
if epoch % 5 == 0 or val_mae < best_val_mae:
print(f"Epoch {epoch:3d} | Loss: {train_loss:.4f} | Val MAE: {val_mae:.4f} "
f"| 1d:{val_mae_h.get('1-day', 0):.4f} "
f"3d:{val_mae_h.get('3-day', 0):.4f} "
f"7d:{val_mae_h.get('7-day', 0):.4f}")
mlflow.log_metrics({
"train_loss": train_loss,
f"val_mae_1d": val_mae_h.get('1-day', float('nan')),
f"val_mae_3d": val_mae_h.get('3-day', float('nan')),
f"val_mae_7d": val_mae_h.get('7-day', float('nan')),
}, step=epoch)
if val_mae < best_val_mae:
best_val_mae = val_mae
best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}
patience_counter = 0
else:
patience_counter += 1
if patience_counter >= PATIENCE:
print(f"\nEarly stopping at epoch {epoch}")
break
# Save
model.load_state_dict(best_state)
torch.save(best_state, MODEL_DIR / 'best_model.pt')
mlflow.log_artifact(MODEL_DIR / 'best_model.pt')
# Beat-baseline check
beat_count = sum(
val_mae_h.get(h, float('inf')) < 0.9 * BASELINE_MAE[h]
for h in ('1-day', '3-day', '7-day')
)
print(f"\nBest Val MAE: {best_val_mae:.4f}")
print(f"Baseline 1d/3d/7d: {BASELINE_MAE['1-day']:.4f}/{BASELINE_MAE['3-day']:.4f}/{BASELINE_MAE['7-day']:.4f}")
print(f"Beats baseline at 0.9x: {beat_count}/3 horizons")
print(f"\nDone! {datetime.now()}")
if __name__ == '__main__':
main()