深度学习基础|第R2周 医疗成本预测

发布时间:2026/10/2 16:20:25
深度学习基础|第R2周 医疗成本预测 第R2周医疗成本预测 本文为365天深度学习训练营 中的学习记录博客 原作者K同学啊编译器jupyterlab一、前期准备1. 数据导入2. 探索热力图numeric_colsdf.select_dtypes(include[int64,float64])plt.figure(figsize(12,8))sns.heatmap(numeric_cols.corr(),cmapcoolwarm,annotTrue)plt.title(Correlation Heatmap)plt.show()Correlation Heatmap相关性热图不同数值变量之间 Pearson 相关系数correlation coefficient的大小。其中insurance_coverage_pct vs annual_medical_cost 的 r-0.86说明二者存在较强烈的线性关系即可能意味着特征冗余。此时纳入机器学习会导致回归参数不稳定、标准误增加。3. 探索分类特征与回归数值的关系3.1 箱线图importmatplotlib.pyplotaspltimportseabornassnsfrommatplotlib.font_managerimportFontProperties#中文字体路径fontFontProperties(fname/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc)# Seaborn风格设置sns.set_style(darkgrid)sns.set(font_scale0.8)# 创建matplotlib的fig对象和子图对象axfig,axplt.subplots(1,3,figsize(12,4))# 多个数值变量的箱线图sns.boxplot(datadf.loc[:,[annual_medical_cost]],axax[0],whis3)ax[0].set_title(多个数值变量,fontpropertiesfont)# 一个数值变量多个分组的箱线图sns.boxplot(xdf[hospital_admissions],ydf[annual_medical_cost],axax[1],whis3)ax[1].set_title(一个数值变量多个分组,fontpropertiesfont)# 一个数值变量多个分组子分组的箱线图sns.boxplot(xhospital_admissions,yannual_medical_cost,huesmoker,datadf,paletteSet1,width0.5,axax[2],whis3)ax[2].set_title(一个数值变量多个分组/子分组,fontpropertiesfont)plt.tight_layout()plt.show()在这里插入图片描述由于字体无法显示原因修改了代码3.2 小提琴图# Seaborn风格设置sns.set(font_scale0.8,styledarkgrid)# 创建fig和子图fig,axplt.subplots(1,3,figsize(12,4))# 多个数值变量的小提琴图sns.violinplot(datadf.loc[:,[annual_medical_cost]],axax[0])ax[0].set_title(多个数值变量,fontpropertiesfont)# 一个数值变量多个分组sns.violinplot(xdf[heart_disease],ydf[annual_medical_cost],axax[1])ax[1].set_title(一个数值变量多个分组,fontpropertiesfont)# 一个数值变量多个分组/子分组sns.violinplot(xheart_disease,yannual_medical_cost,huesmoker,datadf,paletteSet1,width0.5,axax[2])ax[2].set_title(一个数值变量多个分组/子分组,fontpropertiesfont)plt.tight_layout()plt.show()3.3 条形统计图和散点图探索sns.set_style(darkgrid)plt.figure(figsize(6,4))sns.barplot(xheart_disease,yannual_medical_cost,datadf,errorbarci)plt.title(不同心脏病状态的平均医疗费用,fontpropertiesfont)plt.xlabel(Heart Disease)plt.ylabel(Annual Medical Cost)plt.show()plt.figure(figsize(7,4))sns.stripplot(xheart_disease,yannual_medical_cost,datadf,jitterTrue)plt.title(不同心脏病状态下医疗费用分布,fontpropertiesfont)plt.show()4. 探索数值特征与回归特征的关系4.1 气泡图4.2散点图回归线二、数据预处理1. 处理缺失值2. 编码object对象即代表1、3、8、15、17列为类别变量oe OrdinalEncoder() 创建编码器自动分配数字具有大小关系3. 划分训练集与测试集4. 探索字段重要性排行5. 标准化6. 创建dataloaderfromtorch.utils.dataimportDataLoader batch_size32# 封装数据train_datasetdata.TensorDataset(X_train,y_train)test_datasetdata.TensorDataset(X_test,y_test)# 加载数据train_dataloaderDataLoader(train_dataset,batch_sizebatch_size,shuffleTrue)#test_dataloaderDataLoader(test_dataset,batch_sizebatch_size)#, shuffleTrue三、构建模型1. 设置模型参数devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)devicedevice(typecuda)2. 定义模型classmodel_lstm(nn.Module):def__init__(self):super(model_lstm,self).__init__()self.lstm0nn.LSTM(input_size19,hidden_size200,num_layers1,batch_firstTrue)#LSTM内部隐藏状态维度 200维1层self.fc0nn.Linear(200,1)defforward(self,x):out,_self.lstm0(x)outself.fc0(out)returnout modelmodel_lstm()fromtorchinfoimportsummary summary(model,(64,1,19))做了两件事定义一个 LSTM 神经网络模型用 torchinfo.summary() 查看模型结构和参数量3. 编写训练函数deftrain(dataloader,model,loss_fn,optimizer):sizelen(dataloader.dataset)num_batcheslen(dataloader)train_loss,train_acc0,0pred_list[]y_list[]forX,yindataloader:X,yX.to(device),y.to(device)predmodel(X)predpred.squeeze()y_list[i.detach().numpy()foriiny.cpu()]pred_list[i.detach().numpy()foriinpred.cpu()]#lossloss_fn(pred,y)optimizer.zero_grad()loss.backward()optimizer.step()train_lossloss.item()R2metrics.r2_score(y_list,pred_list)#第一个必须是真实值第二个必须是预测值否值 R2 可能会为负数train_loss/num_batchesreturnR2,train_loss4. 编写测试函数deftest(dataloader,model,loss_fn):sizelen(dataloader.dataset)num_batcheslen(dataloader)test_loss,test_acc0,0pred_list[]y_list[]withtorch.no_grad():forX,yindataloader:X,yX.to(device),y.to(device)predmodel(X)predpred.squeeze()y_list[i.detach().numpy()foriiny.cpu()]pred_list[i.detach().numpy()foriinpred.cpu()]lossloss_fn(pred,y)test_lossloss.item()R2metrics.r2_score(y_list,pred_list)test_loss/num_batchesreturnR2,test_loss四、训练模型五、Loss与R2图今天就没有总结啦学习内容都放在各章节里了。这一章对我们临床研究的人很友好顺便还温故了一些统计学知识。