1. 医疗成本预测项目概述
医疗成本预测一直是保险精算和健康管理领域的核心课题。传统线性回归模型在处理这类问题时往往捉襟见肘,因为医疗费用数据通常呈现高度非线性、长尾分布的特征。我在最近一个商业保险项目中尝试使用LSTM网络构建预测模型,相比基线XGBoost模型获得了12%的R2提升。本文将完整还原这个实战案例,重点分享三个关键突破点:如何通过特征工程挖掘医疗数据的时序特性、LSTM超参数调优的实用技巧,以及如何避免医疗数据中常见的过拟合陷阱。
2. 数据准备与探索分析
2.1 数据集特征解析
我们使用的数据集包含10,000名投保人的完整医疗记录,每个样本包含以下特征维度:
- 人口统计学特征:年龄、性别、BMI指数、子女数量
- 行为特征:吸烟状况、区域分布
- 临床指标:血压、血糖等8项体检数据
- 历史费用:过去3年的季度医疗支出记录
- 目标变量:下一年度的总医疗成本
关键发现:通过热力图分析发现,血压指标与医疗费用的相关系数达到0.37,远高于其他临床指标。这提示心血管健康状态可能是医疗成本的重要驱动因素。
2.2 数据可视化洞察
使用Seaborn的pairplot绘制数值特征关系图时,发现两个重要模式:
- BMI与医疗费用呈现明显的"J型"曲线关系,当BMI>30时费用增速显著加快
- 年龄分布呈现双峰特征,需特别注意30-40岁与60-70岁两个人群的行为差异
python复制# 特征关系可视化代码示例
import seaborn as sns
sns.jointplot(x='bmi', y='expenses', data=df, kind='reg',
joint_kws={'line_kws':{'color':'red'}})
3. 特征工程实战
3.1 时序特征构造
原始数据中的季度费用记录本质上是时间序列数据,我们通过以下方法提取时序特征:
- 计算滚动统计量:过去4个季度的移动平均值、标准差
- 构建差分特征:当前季度与去年同期的差值
- 提取趋势特征:使用线性回归拟合过去6个季度的斜率
python复制# 时序特征生成示例
def create_temporal_features(df):
df['rolling_mean'] = df['quarterly_expenses'].rolling(4).mean()
df['year_diff'] = df['quarterly_expenses'] - df['quarterly_expenses'].shift(4)
return df
3.2 分类特征编码
对于吸烟状态等分类变量,测试了三种编码方式:
- One-Hot编码:适合线性模型但增加维度
- Target编码:基于目标变量均值编码,需防范数据泄露
- Embedding编码:通过神经网络学习分布式表示
最终选择Target编码,因其在验证集上比One-Hot提升3%的R2分数,同时保持较低维度。
4. LSTM模型构建
4.1 网络架构设计
我们的BiLSTM模型包含以下核心组件:
- 输入层:接受56维特征向量(原始特征+衍生特征)
- 双向LSTM层:128个单元,dropout=0.2
- 注意力机制层:自动学习特征重要性
- 全连接层:ReLU激活,输出单个预测值
python复制class MedicalCostPredictor(nn.Module):
def __init__(self, input_size):
super().__init__()
self.lstm = nn.LSTM(input_size, 128, bidirectional=True)
self.attention = nn.Sequential(
nn.Linear(256, 128),
nn.Tanh(),
nn.Linear(128, 1, bias=False)
)
self.fc = nn.Linear(256, 1)
def forward(self, x):
lstm_out, _ = self.lstm(x)
attn_weights = F.softmax(self.attention(lstm_out), dim=1)
context = torch.sum(attn_weights * lstm_out, dim=1)
return self.fc(context)
4.2 损失函数优化
针对医疗费用的长尾分布,我们对比了三种损失函数:
- MSE:对异常值敏感但训练稳定
- MAE:更鲁棒但收敛较慢
- Huber Loss:结合二者优点
最终选择Huber Loss,其超参数δ通过网格搜索确定为1.35,在验证集上比MSE提升7%的稳健性。
5. 训练技巧与调优
5.1 学习率调度策略
采用余弦退火配合热重启的策略:
- 初始学习率:3e-4
- 周期长度:20个epoch
- 最小学习率:1e-5
- 重启时学习率倍增系数:0.8
这种设置使得模型在训练后期仍能跳出局部最优,最终测试集R2达到0.87。
5.2 早停机制实现
实现自定义早停策略监控三个指标:
- 验证损失连续5轮不下降
- 训练/验证损失比值超过1.2
- R2分数波动标准差小于0.01
python复制class EarlyStopper:
def __init__(self, patience=5):
self.patience = patience
self.counter = 0
self.min_loss = float('inf')
def should_stop(self, val_loss):
if val_loss < self.min_loss:
self.min_loss = val_loss
self.counter = 0
else:
self.counter += 1
if self.counter >= self.patience:
return True
return False
6. 模型评估与解释
6.1 性能指标对比
在保留测试集上的表现:
| 模型类型 | R2分数 | MAE(美元) | 训练时间 |
|---|---|---|---|
| 线性回归 | 0.62 | 4821 | 2s |
| XGBoost | 0.78 | 3275 | 15m |
| 普通LSTM | 0.83 | 2850 | 1.2h |
| 本文BiLSTM | 0.87 | 2318 | 1.8h |
6.2 误差分析
通过绘制预测值与真实值的残差图,发现模型在高端费用区间(>5万美元)仍有较大误差。这提示我们需要:
- 对高费用样本进行加权处理
- 在高费用区间采用分位数损失
- 收集更多高费用样本增强代表性
7. 部署优化建议
在实际部署时发现两个关键问题:
- 实时性要求:将模型转换为ONNX格式后,推理速度提升3倍
- 内存限制:通过量化将模型大小从420MB压缩到78MB
python复制# 模型量化示例
quantized_model = torch.quantization.quantize_dynamic(
model, {nn.LSTM, nn.Linear}, dtype=torch.qint8
)
8. 经验总结与避坑指南
-
数据泄露预防:在构造时序特征时,必须严格确保只用历史数据计算衍生特征。我们最初错误地使用了未来数据导致线上表现暴跌40%
-
内存优化:医疗数据通常包含大量数值特征,建议:
- 使用float32而非float64
- 对分类变量使用category类型
- 分批生成时序特征
-
超参数敏感度:LSTM的hidden_size对内存消耗呈平方级增长,建议从64开始逐步增加。我们最终选择128作为平衡点
-
解释性增强:通过SHAP值分析发现,血压和既往费用史对预测结果的贡献度超过60%,这为保险产品设计提供了明确方向
