#!/usr/bin/env python3 """ Inference Profiling for Wuhan Respiratory Disease Risk Prediction. Run inference multiple times and report timing statistics. """ 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 time from pathlib import Path from datetime import datetime PROCESSED_DIR = Path('processed') MODEL_DIR = Path('models/spatiotemporal_gcn') 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.""" nodes = pd.read_parquet(PROCESSED_DIR / 'graph' / 'node_features.parquet') return nodes def load_weather_features(n_days=14): """Load weather features for inference.""" lf = pd.read_parquet(PROCESSED_DIR / 'weather' / 'lag_features.parquet') lf['date'] = pd.to_datetime(lf['date']) lf = lf.sort_values('date') 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) return x def prepare_input(nodes, x_weather): """Prepare input tensor for inference.""" elev = nodes['elevation_m'].values pop = nodes['pop_density'].values elev_norm = (elev - elev.mean()) / (elev.std() + 1e-8) spatial_scale = np.clip(1.0 + 0.1 * elev_norm, 0.5, 2.0) x_global = x_weather.numpy() x = np.tile(x_global[np.newaxis, :, :], (len(nodes), 1, 1)) x = x * spatial_scale[:, np.newaxis, np.newaxis] return torch.FloatTensor(x) def run_inference_once(x, edge_index, model_path): """Run single inference and return timing.""" import onnxruntime as ort sess = ort.InferenceSession(model_path, providers=['CPUExecutionProvider']) x_np = x.cpu().numpy().astype(np.float32) edge_np = edge_index.cpu().numpy().astype(np.int64) start = time.perf_counter() risk = sess.run(None, { 'node_features': x_np, 'edge_index': edge_np })[0] elapsed = time.perf_counter() - start return elapsed, risk.shape def profile_inference(model_path=None, n_runs=10, warmup=3): """ Profile inference performance. Args: model_path: Path to ONNX model n_runs: Number of profiling runs warmup: Number of warmup runs (not counted in stats) """ if model_path is None: model_path = MODEL_DIR / 'model_1_3_7.onnx' model_path = Path(model_path) if not model_path.exists(): print(f"Model not found: {model_path}") print("Run train_model.py first to generate the model") return None print(f"\n=== Inference Profiling ===") print(f"Model: {model_path}") print(f"Runs: {n_runs} (+ {warmup} warmup)") # Load data print("\nLoading data...") edge_index = load_graph() n_nodes = edge_index.max().item() + 1 print(f" Graph: {n_nodes} nodes") nodes = load_node_metadata() print(f" Nodes: {len(nodes)}") x_weather = load_weather_features(n_days=14) x = prepare_input(nodes, x_weather) print(f" Input: {x.shape}") # Warmup print(f"\nWarmup ({warmup} runs)...") for i in range(warmup): run_inference_once(x, edge_index, model_path) # Profiling runs print(f"Profiling ({n_runs} runs)...") timings = [] shapes = [] for i in range(n_runs): elapsed, shape = run_inference_once(x, edge_index, model_path) timings.append(elapsed) shapes.append(shape) print(f" Run {i+1}/{n_runs}: {elapsed*1000:.2f} ms") # Statistics timings = np.array(timings) stats = { 'mean_ms': float(timings.mean() * 1000), 'std_ms': float(timings.std() * 1000), 'min_ms': float(timings.min() * 1000), 'max_ms': float(timings.max() * 1000), 'median_ms': float(np.median(timings) * 1000), 'p95_ms': float(np.percentile(timings, 95) * 1000), 'p99_ms': float(np.percentile(timings, 99) * 1000), 'runs': n_runs, 'model_path': str(model_path), 'n_nodes': n_nodes, 'input_shape': list(x.shape), 'output_shape': list(shapes[0]), 'timestamp': datetime.now().isoformat() } # Report print(f"\n=== Performance Summary ===") print(f" Mean: {stats['mean_ms']:.2f} ms") print(f" Std: {stats['std_ms']:.2f} ms") print(f" Min: {stats['min_ms']:.2f} ms") print(f" Max: {stats['max_ms']:.2f} ms") print(f" Median: {stats['median_ms']:.2f} ms") print(f" P95: {stats['p95_ms']:.2f} ms") print(f" P99: {stats['p99_ms']:.2f} ms") print(f"\n Throughput: {1000/stats['mean_ms']:.1f} inferences/sec") print(f" Daily batch (365 runs): {stats['mean_ms']*365/1000:.2f} sec/day") # Save profile results output_dir = Path('outputs/profiles') output_dir.mkdir(exist_ok=True) output_file = output_dir / f'profile_{datetime.now().strftime("%Y%m%d_%H%M%S")}.json' import json with open(output_file, 'w') as f: json.dump(stats, f, indent=2) print(f"\nSaved: {output_file}") return stats if __name__ == '__main__': import argparse parser = argparse.ArgumentParser(description='Profile inference performance') parser.add_argument('--model', type=str, default=None, help='Path to ONNX model') parser.add_argument('--runs', type=int, default=10, help='Number of profiling runs') parser.add_argument('--warmup', type=int, default=3, help='Number of warmup runs') args = parser.parse_args() profile_inference(args.model, args.runs, args.warmup)