1. 项目概述:基于机器学习的COVID-19病例预测
这个项目来自台大2021年春季机器学习课程的第一个作业,目标是利用机器学习技术预测COVID-19确诊病例数。作为疫情初期最具挑战性的预测问题之一,它考验着我们对时序数据处理、特征工程和回归模型的综合掌握能力。
在实际操作中,我们需要处理来自多个国家和地区的不完整疫情数据,构建有效的特征表示,并选择合适的机器学习模型进行训练和预测。这个项目不仅具有学术意义,其方法论也能直接应用于其他传染病预测、经济指标预测等现实场景。我完成这个作业后,将核心思路和实操经验整理成文,特别适合刚入门机器学习、想通过实战项目提升技能的朋友参考。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心问题分析与数据准备
2.1 预测任务定义
我们需要预测的是未来n天(作业中n=30)的每日新增确诊病例数。这本质上是一个时间序列回归问题,但与传统时序预测不同,疫情数据具有几个显著特点:
- 数据质量参差不齐:不同国家/地区的检测能力、报告标准不一致
- 存在明显的政策干预影响:封城、社交限制等措施会突然改变传播趋势
- 具有空间传播特性:周边国家的疫情发展会相互影响
2.2 数据获取与清洗
原始数据来自约翰霍普金斯大学的COVID-19数据集,包含以下关键字段:
- 日期(date):从2020-01-22开始的时间序列
- 国家/地区(country/region)
- 省/州(province/state):部分国家的细分区域数据
- 确诊数(confirmed):累计确诊病例
- 死亡数(deaths):累计死亡病例
- 康复数(recovered):累计康复病例
数据清洗的关键步骤:
python复制# 示例数据清洗代码
def clean_data(df):
# 处理缺失值:用前向填充+国家均值填充
df['confirmed'] = df.groupby('country')['confirmed'].apply(
lambda x: x.fillna(method='ffill').fillna(x.mean()))
# 计算每日新增而非累计值
df['new_cases'] = df.groupby('country')['confirmed'].diff().fillna(0)
# 去除异常值:超过3倍标准差的值用移动平均替代
rolling_mean = df.groupby('country')['new_cases'].transform(
lambda x: x.rolling(7, min_periods=1).mean())
std = df.groupby('country')['new_cases'].std()
df.loc[df['new_cases'] > 3*std, 'new_cases'] = rolling_mean
return df
注意:实际作业中还需要处理国家/地区名称不一致(如"US" vs "United States")、数据上报延迟等问题。建议建立专门的国家名称映射表。
2.3 特征工程策略
有效的特征工程是预测准确的关键。我们设计了以下几类特征:
-
时间特征:
- 星期几(疫情传播往往有周周期性)
- 是否为节假日(影响人群聚集程度)
- 距离首例确诊的天数(反映疫情发展阶段)
-
统计特征:
- 过去3/7/14天的移动平均值
- 过去3/7/14天的变化率
- 累计确诊数的对数(反映疫情规模)
-
外部特征:
- 政府响应严格指数(来自牛津COVID-19政府响应追踪器)
- 周边国家/地区的疫情情况
- 疫苗接种进度(后期数据)
-
衍生特征:
- 传播速率(Rt值)的估计
- 疫情波次识别(通过峰值检测)
python复制# 特征生成示例
def create_features(df):
# 移动平均特征
for window in [3,7,14]:
df[f'ma_{window}'] = df.groupby('country')['new_cases'].transform(
lambda x: x.rolling(window).mean())
# 变化率特征
df['change_rate_7d'] = df['ma_7'] / df.groupby('country')['ma_7'].shift(7)
# 星期特征
df['day_of_week'] = pd.to_datetime(df['date']).dt.dayofweek
return df
3. 模型选择与实现
3.1 模型选型比较
我们对比了几种适合时序预测的机器学习模型:
| 模型类型 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 线性回归 | 简单、可解释性强 | 无法捕捉非线性关系 | 趋势稳定的短期预测 |
| 随机森林 | 处理非线性关系、抗噪声 | 难以捕捉长期依赖 | 中等复杂度数据 |
| XGBoost | 表现稳定、特征重要性 | 需要调参 | 大多数场景 |
| LSTM | 擅长时序模式识别 | 需要大量数据、训练慢 | 复杂长期依赖 |
| 集成模型 | 结合各模型优势 | 复杂度高 | 追求最高精度 |
基于作业数据量和复杂度,我最终选择了XGBoost作为基础模型,原因如下:
- 能够自动处理特征间的非线性关系
- 提供特征重要性评估,便于分析关键影响因素
- 对缺失值相对鲁棒
- 训练速度较快,适合作业时间要求
3.2 XGBoost实现细节
核心参数设置与调优策略:
python复制import xgboost as xgb
params = {
'objective': 'reg:squarederror',
'learning_rate': 0.05,
'max_depth': 6,
'subsample': 0.8,
'colsample_bytree': 0.8,
'n_estimators': 1000,
'early_stopping_rounds': 50,
'eval_metric': 'rmse'
}
# 时间序列交叉验证
tss = TimeSeriesSplit(n_splits=5)
for train_idx, val_idx in tss.split(X):
X_train, X_val = X.iloc[train_idx], X.iloc[val_idx]
y_train, y_val = y.iloc[train_idx], y.iloc[val_idx]
model = xgb.XGBRegressor(**params)
model.fit(X_train, y_train,
eval_set=[(X_val, y_val)],
verbose=False)
# 记录每次验证结果...
提示:对于时间序列数据,绝对不能使用随机交叉验证,必须按时间顺序划分训练/验证集,否则会导致数据泄露(未来信息污染过去预测)。
3.3 模型集成策略
为进一步提升预测稳定性,我实现了两种集成方法:
-
多模型集成:
- 训练XGBoost、LightGBM和随机森林三个模型
- 用简单平均或线性回归学习各模型的权重
-
时序集成:
- 用滑动窗口训练多个模型(如每月重新训练一次)
- 预测时组合最近几个窗口模型的输出
python复制# 多模型集成示例
from sklearn.ensemble import RandomForestRegressor
from lightgbm import LGBMRegressor
models = {
'xgb': xgb.XGBRegressor(**xgb_params),
'lgb': LGBMRegressor(**lgb_params),
'rf': RandomForestRegressor(**rf_params)
}
# 训练各模型
for name, model in models.items():
model.fit(X_train, y_train)
# 集成预测
def ensemble_predict(X):
preds = [model.predict(X) for model in models.values()]
return np.mean(preds, axis=0)
4. 评估与优化
4.1 评估指标选择
不同于一般回归问题,疫情预测需要特别关注以下指标:
-
RMSE(均方根误差):
- 作业主要评估指标
- 对异常值敏感,反映整体偏差
-
MAE(平均绝对误差):
- 更直观的解释(平均每天差多少例)
- 对异常值不敏感
-
MAPE(平均绝对百分比误差):
- 相对误差,适合比较不同规模的地区
- 当真实值接近0时不稳定
-
趋势准确率:
- 预测方向(上升/下降)的正确率
- 对政策制定更有参考价值
4.2 误差分析与改进
通过分析验证集上的预测错误,发现几个常见问题及解决方案:
| 问题现象 | 可能原因 | 解决方案 |
|---|---|---|
| 持续低估峰值 | 模型对突变不敏感 | 添加变化率阈值作为特征 |
| 节假日预测差 | 未考虑特殊日期影响 | 加入节假日日历特征 |
| 小国家误差大 | 数据量不足 | 按地区聚类或使用迁移学习 |
| 长期预测发散 | 误差累积效应 | 采用滚动预测而非直接多步预测 |
改进后的特征重要性分析示例(XGBoost输出):
code复制特征 重要性
ma_7 0.32
government_response 0.18
day_of_week 0.12
change_rate_7d 0.09
neighbor_cases 0.07
... ...
4.3 最终模型表现
经过多轮优化,模型在测试集上的表现:
| 国家 | RMSE | MAE | MAPE |
|---|---|---|---|
| 美国 | 4231 | 2987 | 15.2% |
| 英国 | 1872 | 1345 | 18.7% |
| 日本 | 892 | 643 | 22.3% |
| 巴西 | 2543 | 1821 | 17.9% |
| 全球平均 | 2145 | 1532 | 19.1% |
注意:这些数字是示例,实际作业结果会根据数据预处理和模型选择的差异而变化。小国家通常误差更大,因为训练数据较少。
5. 实际应用与扩展
5.1 部署为预测服务
将训练好的模型部署为API服务的核心步骤:
- 模型持久化:
python复制import joblib
joblib.dump(model, 'covid_predictor.pkl')
- 构建Flask API:
python复制from flask import Flask, request, jsonify
import pandas as pd
app = Flask(__name__)
model = joblib.load('covid_predictor.pkl')
@app.route('/predict', methods=['POST'])
def predict():
data = request.json
df = pd.DataFrame(data)
features = preprocess(df)
preds = model.predict(features)
return jsonify({'predictions': preds.tolist()})
- 添加自动更新:
- 设置定期任务(如每周)重新训练模型
- 当新数据超过一定阈值时触发重新训练
5.2 扩展应用方向
这个项目的技术栈可以扩展到许多相关领域:
-
其他传染病预测:
- 流感季节预测
- 登革热爆发预警
-
社会经济指标预测:
- 失业率变化
- 零售销售额预测
-
商业应用:
- 产品需求预测
- 服务器流量预测
5.3 后续改进思路
如果想进一步提升模型性能,可以考虑:
-
引入更多外部数据:
- 移动设备的位置数据(反映人员流动)
- 航空客运量数据
- 社交媒体舆情分析
-
改进模型架构:
- 尝试Transformer-based时序模型
- 结合SEIR等流行病学模型
-
不确定性量化:
- 输出预测区间而不仅是点估计
- 使用分位数回归或贝叶斯方法
完成这个项目后,我最大的体会是:真实世界的数据远比教科书上的示例复杂,优秀的机器学习工程师需要同时具备数据处理、领域知识和模型调优的能力。特别是在疫情预测这种具有重大社会影响的应用中,我们需要对模型的局限性保持清醒认识,避免过度自信的预测。
