1. 项目背景与核心目标
这个机器学习实战项目源自李宏毅教授2022年春季《机器学习》课程的第一次作业,要求基于历史数据构建新冠病毒感染人数的预测模型。作为入门级的时序预测任务,它完美融合了公共卫生事件分析与基础机器学习技术应用两大维度。
在实际操作中,我们需要处理来自真实世界的非平稳时序数据,构建能够捕捉感染人数变化规律的预测模型。这个项目的独特价值在于:
- 数据维度真实:使用经过脱敏处理的实际感染统计数据
- 问题定义明确:单变量时序预测任务(Univariate Time Series Forecasting)
- 评估标准清晰:以均方误差(MSE)作为核心指标
提示:虽然作业本身使用2021年的数据,但方法论完全适用于当前各类疫情数据分析场景,包括流感等其他传染病的预测。
2. 数据理解与特征工程
2.1 原始数据解析
原始数据集通常包含两列核心信息:
- 日期(date):连续的时间戳
- 确诊数(cases):当日新增感染人数
通过EDA分析可以发现几个关键特征:
- 明显的周期性波动(通常以7天为周期)
- 存在异常峰值(节假日或检测策略变化导致)
- 非平稳性趋势(感染波峰波谷差异显著)
# 典型的数据加载代码示例 import pandas as pd data = pd.read_csv('covid_cases.csv', parse_dates=['date']) print(data.describe())2.2 关键特征构建
基于时序预测的常用方法,我们需要构造以下特征类型:
| 特征类型 | 生成方法 | 作用说明 |
|---|---|---|
| 滞后特征 | 前1/7/14天的病例数 | 捕捉短期依赖关系 |
| 移动统计 | 7天平均/标准差 | 平滑噪声反映趋势 |
| 时间特征 | 星期几/月份/季度 | 捕获周期性模式 |
| 变化率 | 日环比/周同比 | 反映增长加速度 |
# 特征工程示例代码 data['lag_1'] = data['cases'].shift(1) data['rolling_7_mean'] = data['cases'].rolling(7).mean() data['day_of_week'] = data['date'].dt.dayofweek3. 模型构建与技术选型
3.1 基线模型选择
作业中通常会对比三类经典方法:
简单移动平均:
- 实现简单但效果有限
- 适合建立评估基准
def moving_average(data, window=7): return data.rolling(window).mean()线性回归:
- 使用前述构造的特征
- 可解释性强但非线性关系捕捉有限
神经网络模型:
- 全连接网络(FCN)作为基础架构
- 输入层节点数对应特征维度
- 隐藏层通常2-3层即可
3.2 深度模型优化技巧
对于神经网络的实现,有几个关键优化点:
数据标准化:
from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X_train = scaler.fit_transform(X_train)损失函数选择:
- MSE直接对应评估指标
- 可尝试SmoothL1Loss减少异常值影响
早停机制:
from pytorch_lightning.callbacks import EarlyStopping early_stop = EarlyStopping(monitor='val_loss', patience=5)
4. 模型评估与结果分析
4.1 评估指标解读
使用MSE作为核心指标时需注意:
- 对异常值敏感
- 量级与数据规模相关
- 建议同时观察MAE和MAPE
注意:在公共卫生领域,过高预测和过低预测的风险不对称,可能需要设计非对称损失函数。
4.2 典型结果模式
通过多次实验通常会观察到:
| 模型类型 | 验证集MSE | 过拟合风险 | 训练速度 |
|---|---|---|---|
| 移动平均 | 较高 | 低 | 极快 |
| 线性回归 | 中等 | 中 | 快 |
| 神经网络 | 较低 | 高 | 慢 |
4.3 可视化分析技巧
建议使用以下可视化方法诊断模型:
预测-实际对比图:
plt.plot(y_true, label='Actual') plt.plot(y_pred, label='Predicted')残差分布图:
sns.distplot(y_true - y_pred)滚动误差图:
plt.plot(moving_average(np.abs(y_true - y_pred)))
5. 实战经验与避坑指南
5.1 数据预处理陷阱
缺失值处理:
- 直接填充0会引入偏差
- 建议使用前后均值或插值法
数据泄露:
- 移动统计量计算时需严格区分训练/测试集
- 使用
TimeSeriesSplit进行交叉验证
5.2 模型训练技巧
批次大小选择:
- 小批次(32-64)更适合时序数据
- 太大容易错过局部波动模式
学习率设置:
- 初始建议1e-3到1e-4
- 使用学习率调度器(如ReduceLROnPlateau)
正则化策略:
- L2正则系数建议0.01-0.001
- Dropout率建议0.2-0.5
5.3 部署注意事项
模型更新频率:
- 建议每周重新训练
- 保留历史模型用于比对
预测不确定性:
- 输出预测区间而非单点估计
- 可使用MC Dropout估算方差
业务解释性:
- 提供特征重要性分析
- 生成趋势分解图表
6. 项目扩展方向
对于希望深入研究的同学,可以考虑以下进阶方向:
多变量时序模型:
- 加入疫苗接种率、防控政策等外部变量
- 使用LSTM/Transformer架构
空间维度扩展:
- 构建地区级预测模型
- 加入地理邻接矩阵
实时预测系统:
# 简易API示例 from fastapi import FastAPI app = FastAPI() @app.post("/predict") async def predict(date: str): return {"prediction": model.predict(date)}异常检测集成:
- 自动识别数据上报异常
- 结合统计检验方法
这个项目虽然作为课程作业出现,但完整覆盖了从数据预处理到模型部署的机器学习全流程。在实际操作中,我发现时序数据的季节性分解质量会显著影响最终效果,建议使用STL分解而非传统方法。另外,对于突发的感染高峰,单纯的统计学习可能表现不佳,这时需要结合流行病学领域的专业知识进行模型校正。