Files
CA/scripts/evaluate.py

521 lines
20 KiB
Python
Raw Permalink Normal View History

#!/usr/bin/env python3
"""
Model Evaluation Script - Phase 3.8
Evaluates trained Spatial-Temporal GCN model on held-out test data (December 2023).
Generates comprehensive markdown report with per-horizon MAE, risk classification analysis,
and baseline comparison.
Test Period: 2023-12-01 to 2023-12-31 (not used in training/validation)
"""
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
from pathlib import Path
from datetime import datetime
from sklearn.metrics import accuracy_score, precision_recall_fscore_support, confusion_matrix
import json
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')
REPORTS_DIR = Path('reports')
REPORTS_DIR.mkdir(exist_ok=True)
TEST_START = '2023-12-01'
TEST_END = '2023-12-31'
BASELINE_MAE = {'1-day': 0.2314, '3-day': 0.5424, '7-day': 0.6391}
RISK_THRESHOLDS = {
'low': 0.33,
'medium': 0.66,
'high': 1.0
}
def load_test_data():
"""Load test data for December 2023."""
print("Loading test data...")
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
)
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")
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")
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)
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.
"""
n_nodes = len(nodes)
n_days = len(dates)
global_risk = medical.groupby('date')['risk'].mean()
global_risk_dict = global_risk.to_dict()
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]
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."""
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)
elev_scale = 1.0 + 0.1 * elev_norm
elev_scale = np.clip(elev_scale, 0.5, 2.0).astype(np.float32)
pop_scale = np.ones_like(elev_scale)
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."""
batch_size = len(node_indices)
T, F = x_global.shape
batch_elev = elev_scale[node_indices]
x = np.tile(x_global[np.newaxis, :, :], (batch_size, 1, 1))
x = x * batch_elev[:, np.newaxis, np.newaxis]
return x.astype(np.float32)
def evaluate_model(model, x_global, elev_scale, y, edge_index, window=14, batch_size=512):
"""
Comprehensive evaluation with per-horizon predictions.
Returns:
results: dict with per-horizon MAE, RMSE,
all_preds: dict with predictions per horizon
all_actuals: dict with actual values per horizon
"""
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 = {}
all_preds = {h: [] for h in horizons}
all_actuals = {h: [] for h in horizons}
print(f"\nEvaluating on {T - window + 1} time windows...")
with torch.no_grad():
for name, h in horizons.items():
if h > T - window:
results[name] = {'mae': float('nan'), 'rmse': float('nan'), 'r2': float('nan')}
continue
preds_list = []
actuals_list = []
for t in range(window, T - h + 1):
for node_start in range(0, n_nodes, batch_size):
node_end = min(node_start + batch_size, n_nodes)
node_indices = np.arange(node_start, node_end)
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)
sub_edge_index, _ = subgraph(node_indices_torch, edge_index, relabel_nodes=False)
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)
valid_mask = ~torch.isnan(y_actual)
if valid_mask.sum() == 0:
continue
pred = model(x_win, sub_edge_index)[valid_mask, :]
horizon_idx = {'1-day': 0, '3-day': 1, '7-day': 2}[name]
preds_list.append(pred[:, horizon_idx].cpu().numpy())
actuals_list.append(y_actual[valid_mask].cpu().numpy())
if preds_list:
preds = np.concatenate(preds_list)
actuals = np.concatenate(actuals_list)
mae = np.mean(np.abs(preds - actuals))
rmse = np.sqrt(np.mean((preds - actuals) ** 2))
ss_res = np.sum((actuals - preds) ** 2)
ss_tot = np.sum((actuals - np.mean(actuals)) ** 2)
r2 = 1 - (ss_res / (ss_tot + 1e-8))
results[name] = {
'mae': float(mae),
'rmse': float(rmse),
'r2': float(r2),
'n_samples': len(preds)
}
all_preds[name] = preds
all_actuals[name] = actuals
print(f" {name}: MAE={mae:.4f}, RMSE={rmse:.4f}, R²={r2:.4f} (n={len(preds)})")
else:
results[name] = {'mae': float('nan'), 'rmse': float('nan'), 'r2': float('nan')}
return results, all_preds, all_actuals
def analyze_risk_classification(all_preds, all_actuals):
"""Analyze risk level classification performance."""
print("\nAnalyzing risk classification...")
results = {}
for horizon in ['1-day', '3-day', '7-day']:
if horizon not in all_preds or len(all_preds[horizon]) == 0:
continue
preds = all_preds[horizon]
actuals = all_actuals[horizon]
def to_category(values):
cats = np.zeros(len(values), dtype=int)
cats[values < RISK_THRESHOLDS['low']] = 0
cats[(values >= RISK_THRESHOLDS['low']) & (values < RISK_THRESHOLDS['medium'])] = 1
cats[values >= RISK_THRESHOLDS['medium']] = 2
return cats
pred_cats = to_category(preds)
actual_cats = to_category(actuals)
accuracy = accuracy_score(actual_cats, pred_cats)
precision, recall, f1, _ = precision_recall_fscore_support(
actual_cats, pred_cats, average='weighted', zero_division=0
)
cm = confusion_matrix(actual_cats, pred_cats, labels=[0, 1, 2])
results[horizon] = {
'accuracy': float(accuracy),
'precision': float(precision),
'recall': float(recall),
'f1': float(f1),
'confusion_matrix': cm.tolist(),
'category_names': ['Low', 'Medium', 'High']
}
print(f" {horizon}: Accuracy={accuracy:.3f}, F1={f1:.3f}")
return results
def generate_report(eval_results, classification_results, model_params, output_path):
"""Generate comprehensive markdown report."""
report = f"""# Model Evaluation Report - Phase 3.8
**Generated:** {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}
**Test Period:** {TEST_START} to {TEST_END}
**Model:** Spatial-Temporal GCN (Transformer + Graph Convolution)
---
## Executive Summary
This report evaluates the trained Spatial-Temporal GCN model on held-out test data (December 2023),
which was not used during training or validation. The model predicts respiratory disease risk at
three forecasting horizons: 1-day, 3-day, and 7-day ahead.
### Key Findings
| Metric | 1-Day Horizon | 3-Day Horizon | 7-Day Horizon |
|--------|---------------|---------------|---------------|
| **MAE** | {eval_results.get('1-day', {}).get('mae', 'N/A'):.4f} | {eval_results.get('3-day', {}).get('mae', 'N/A'):.4f} | {eval_results.get('7-day', {}).get('mae', 'N/A'):.4f} |
| **RMSE** | {eval_results.get('1-day', {}).get('rmse', 'N/A'):.4f} | {eval_results.get('3-day', {}).get('rmse', 'N/A'):.4f} | {eval_results.get('7-day', {}).get('rmse', 'N/A'):.4f} |
| **** | {eval_results.get('1-day', {}).get('r2', 'N/A'):.4f} | {eval_results.get('3-day', {}).get('r2', 'N/A'):.4f} | {eval_results.get('7-day', {}).get('r2', 'N/A'):.4f} |
| **Samples** | {eval_results.get('1-day', {}).get('n_samples', 'N/A')} | {eval_results.get('3-day', {}).get('n_samples', 'N/A')} | {eval_results.get('7-day', {}).get('n_samples', 'N/A')} |
### Baseline Comparison
| Horizon | Baseline MAE | Model MAE | Improvement | Beats 0.9× Baseline? |
|---------|--------------|-----------|-------------|----------------------|
| 1-Day | {BASELINE_MAE['1-day']:.4f} | {eval_results.get('1-day', {}).get('mae', float('inf')):.4f} | {((BASELINE_MAE['1-day'] - eval_results.get('1-day', {}).get('mae', 0)) / BASELINE_MAE['1-day'] * 100):.1f}% | {'✅ Yes' if eval_results.get('1-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['1-day'] else '❌ No'} |
| 3-Day | {BASELINE_MAE['3-day']:.4f} | {eval_results.get('3-day', {}).get('mae', float('inf')):.4f} | {((BASELINE_MAE['3-day'] - eval_results.get('3-day', {}).get('mae', 0)) / BASELINE_MAE['3-day'] * 100):.1f}% | {'✅ Yes' if eval_results.get('3-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['3-day'] else '❌ No'} |
| 7-Day | {BASELINE_MAE['7-day']:.4f} | {eval_results.get('7-day', {}).get('mae', float('inf')):.4f} | {((BASELINE_MAE['7-day'] - eval_results.get('7-day', {}).get('mae', 0)) / BASELINE_MAE['7-day'] * 100):.1f}% | {'✅ Yes' if eval_results.get('7-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['7-day'] else '❌ No'} |
---
## Model Architecture
| Component | Configuration |
|-----------|---------------|
| **Node Features** | {model_params.get('node_features', 48)} (48 weather variables) |
| **Temporal Encoder** | Transformer ({model_params.get('temporal_layers', 3)} layers, {model_params.get('temporal_heads', 4)} heads) |
| **GCN Layers** | [{model_params.get('node_features', 48)} {model_params.get('gcn_hidden', 128)} {model_params.get('gcn_output', 64)}] |
| **Output** | 3 risk horizons (1-day, 3-day, 7-day) |
| **Total Parameters** | {model_params.get('total_params', 'N/A'):,} |
| **Input Window** | {model_params.get('window', 14)} days |
---
## Detailed Evaluation Metrics
### 1-Day Horizon
- **MAE:** {eval_results.get('1-day', {}).get('mae', 'N/A'):.4f}
- **RMSE:** {eval_results.get('1-day', {}).get('rmse', 'N/A'):.4f}
- **:** {eval_results.get('1-day', {}).get('r2', 'N/A'):.4f}
- **Valid Samples:** {eval_results.get('1-day', {}).get('n_samples', 'N/A')}
#### Risk Classification Performance
"""
for horizon in ['1-day', '3-day', '7-day']:
if horizon in classification_results:
cls = classification_results[horizon]
report += f"""
### {horizon} Risk Classification
- **Accuracy:** {cls['accuracy']:.3f}
- **Precision (weighted):** {cls['precision']:.3f}
- **Recall (weighted):** {cls['recall']:.3f}
- **F1 Score (weighted):** {cls['f1']:.3f}
#### Confusion Matrix
| Actual \\ Predicted | Low | Medium | High |
|---------------------|-----|--------|------|
| **Low** | {cls['confusion_matrix'][0][0]} | {cls['confusion_matrix'][0][1]} | {cls['confusion_matrix'][0][2]} |
| **Medium** | {cls['confusion_matrix'][1][0]} | {cls['confusion_matrix'][1][1]} | {cls['confusion_matrix'][1][2]} |
| **High** | {cls['confusion_matrix'][2][0]} | {cls['confusion_matrix'][2][1]} | {cls['confusion_matrix'][2][2]} |
"""
beat_count = sum(
eval_results.get(h, {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE[h]
for h in ['1-day', '3-day', '7-day']
)
report += f"""---
## Conclusions
### Acceptance Criteria Assessment
**Primary Criterion:** Model MAE must be < 0.9 × Baseline MAE for at least one horizon.
**Result:** {'✅ PASSED' if beat_count >= 1 else '❌ FAILED'} ({beat_count}/3 horizons beat baseline at 0.9× threshold)
### Observations
1. **Short-term prediction (1-day):** {'Strong performance with MAE significantly below baseline.' if eval_results.get('1-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['1-day'] else 'Moderate performance, room for improvement.'}
2. **Medium-term prediction (3-day):** {'Good generalization to 3-day horizon.' if eval_results.get('3-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['3-day'] else 'Performance degrades as expected with longer horizon.'}
3. **Long-term prediction (7-day):** {'Excellent 7-day forecasting capability.' if eval_results.get('7-day', {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE['7-day'] else 'Expected challenge with 7-day horizon due to weather prediction uncertainty.'}
### Recommendations for Phase 4
1. **Feature Engineering:** Consider adding additional spatial features (land use, traffic patterns)
2. **Temporal Dynamics:** Experiment with longer input windows (21-30 days)
3. **Model Architecture:** Explore graph attention networks (GAT) for adaptive spatial weighting
4. **Ensemble Methods:** Combine multiple model runs for uncertainty quantification
5. **Real-time Validation:** Implement continuous monitoring on incoming data
---
## Technical Details
### Data Preprocessing
- **Weather Features:** 48 variables (15 pollutant types × 24h + derived features)
- **Spatial Features:** Elevation, population density (used for node-level scaling)
- **Target Variable:** District-level medical risk (weighted outpatient + inpatient cases)
- **Normalization:** Per-node z-score normalization
### Evaluation Methodology
- **Test Set:** December 2023 (completely held out from training/validation)
- **Batch Size:** 512 nodes per batch (memory-efficient evaluation)
- **Metrics:** MAE, RMSE, for regression; Accuracy, F1 for classification
- **Risk Thresholds:** Low (<0.33), Medium (0.33-0.66), High (>0.66)
### Reproducibility
- **Model Checkpoint:** `models/spatiotemporal_gcn/best_model.pt`
- **Evaluation Script:** `scripts/evaluate.py`
- **Random Seed:** 42 (consistent with training)
---
*Report generated by Wuhan Respiratory Disease Risk Prediction System*
"""
with open(output_path, 'w', encoding='utf-8') as f:
f.write(report)
print(f"\nReport saved to: {output_path}")
def main():
print(f"\n{'='*60}")
print(f"Model Evaluation - Phase 3.8 {datetime.now()}")
print(f"{'='*60}")
edge_index, nodes, lf, medical = load_test_data()
n_nodes = len(nodes)
x_global, weather_dates = build_global_weather_timeseries(lf)
targets, _ = build_node_targets(nodes, medical, weather_dates)
elev_scale, pop_scale = build_spatial_scalars(nodes)
dates_arr = pd.to_datetime(weather_dates)
test_mask = (dates_arr >= TEST_START) & (dates_arr <= TEST_END)
x_global_test = x_global[test_mask]
y_test = targets[:, test_mask]
test_days = len(x_global_test)
print(f"\nTest period: {TEST_START} to {TEST_END}")
print(f"Test samples: {test_days} days")
print(f"Global weather shape: {x_global_test.shape}")
print(f"Target shape: {y_test.shape}")
model_path = MODEL_DIR / 'best_model.pt'
if not model_path.exists():
print(f"\n❌ ERROR: Model checkpoint not found at {model_path}")
print("Please run scripts/train_model.py first.")
sys.exit(1)
print(f"\nLoading model from: {model_path}")
model = SpatialTemporalGCN(
node_features=48,
temporal_heads=4,
temporal_layers=3,
gcn_hidden=128,
gcn_output=64,
dropout=0.2
).to(DEVICE)
state_dict = torch.load(model_path, map_location=DEVICE)
model.load_state_dict(state_dict)
model.eval()
total_params = sum(p.numel() for p in model.parameters())
print(f"Model parameters: {total_params:,}")
WINDOW = 14
eval_results, all_preds, all_actuals = evaluate_model(
model, x_global_test, elev_scale, y_test, edge_index,
window=WINDOW, batch_size=512
)
classification_results = analyze_risk_classification(all_preds, all_actuals)
model_params = {
'node_features': 48,
'temporal_heads': 4,
'temporal_layers': 3,
'gcn_hidden': 128,
'gcn_output': 64,
'window': WINDOW,
'total_params': total_params
}
report_path = REPORTS_DIR / 'model_evaluation_phase3.md'
generate_report(eval_results, classification_results, model_params, report_path)
print(f"\n{'='*60}")
print("EVALUATION SUMMARY")
print(f"{'='*60}")
beat_count = sum(
eval_results.get(h, {}).get('mae', float('inf')) < 0.9 * BASELINE_MAE[h]
for h in ['1-day', '3-day', '7-day']
)
for horizon in ['1-day', '3-day', '7-day']:
mae = eval_results.get(horizon, {}).get('mae', float('nan'))
baseline = BASELINE_MAE[horizon]
improvement = ((baseline - mae) / baseline * 100) if not np.isnan(mae) else 0
beats = '' if mae < 0.9 * baseline else ''
print(f"{horizon}: MAE={mae:.4f} (Baseline: {baseline:.4f}, Improvement: {improvement:+.1f}%) {beats}")
print(f"\nAcceptance Criteria: {'✅ PASSED' if beat_count >= 1 else '❌ FAILED'} ({beat_count}/3 horizons)")
print(f"\nFull report: {report_path}")
print(f"{'='*60}\n")
if __name__ == '__main__':
main()