随机森林实战:泰坦尼克号生存预测全流程解析
1. 项目概述:当机器学习遇见历史数据
泰坦尼克号生存预测是数据科学领域的经典入门项目,相当于机器学习界的"Hello World"。这个数据集包含了891名乘客的真实信息,包括年龄、性别、舱位等级等特征,以及最重要的生存状态标签。通过构建随机森林模型,我们不仅能学习到机器学习的基本流程,更能理解特征工程如何影响预测效果。
我在金融风控领域使用随机森林多年,发现这个算法对数据质量的要求相对宽容,非常适合新手作为第一个实战项目。相比逻辑回归等线性模型,随机森林能自动捕捉特征间的非线性关系,而相比深度学习又不需要复杂调参。下面我会分享从数据清洗到模型优化的完整过程,包括几个教科书上不会写的实战技巧。
2. 环境准备与数据加载
2.1 基础工具栈配置
推荐使用Python 3.8+环境,主要依赖库包括:
- pandas 1.3+(数据处理)
- scikit-learn 1.0+(机器学习)
- matplotlib 3.5+(可视化)
# 安装命令(已安装可跳过)
pip install pandas scikit-learn matplotlib
2.2 数据加载与初探
数据集可以从Kaggle官网获取,包含三个关键文件:
- train.csv (训练集,含生存标签)
- test.csv (测试集,不含标签)
- gender_submission.csv (提交样例)
import pandas as pd
train_df = pd.read_csv('train.csv')
test_df = pd.read_csv('test.csv')
print(f"训练集形状: {train_df.shape}")
print(f"测试集形状: {test_df.shape}")
train_df.head()
注意:实际路径需根据文件存放位置调整。首次运行时建议先检查各列数据类型和缺失值情况。
3. 特征工程实战技巧
3.1 缺失值处理的行业经验
年龄(Age)列约有20%缺失值,常见处理方式包括:
- 直接删除(会损失样本)
- 用均值/中位数填充(可能引入偏差)
- 构建预测模型估算(最合理但复杂)
这里演示中位数填充的优化写法:
def fill_na_median(df, column):
"""带异常值保护的填充函数"""
try:
median = df[column].median()
df[column] = df[column].fillna(median)
return df
except KeyError:
print(f"警告:列 {column} 不存在")
return df
train_df = fill_na_median(train_df, 'Age')
3.2 特征构造的创造性思维
原始特征往往需要组合才能发挥价值:
- 家庭规模 = SibSp(兄弟姐妹) + Parch(父母子女)
- 是否独自乘船 = (家庭规模 == 0)
- 姓名中的称谓提取(反映社会地位)
# 提取姓名中的称谓
train_df['Title'] = train_df['Name'].str.extract(' ([A-Za-z]+)\.', expand=False)
title_mapping = {
"Mr": 1, "Miss": 2, "Mrs": 3,
"Master": 4, "Dr": 5, "Rev": 6,
"Col": 7, "Major": 7, "Mlle": 2,
"Countess": 3, "Ms": 2, "Lady": 3,
"Jonkheer": 1, "Don": 1, "Dona": 3,
"Mme": 3, "Capt": 7, "Sir": 1
}
train_df['Title'] = train_df['Title'].map(title_mapping)
4. 随机森林模型构建
4.1 基础模型搭建
先划分特征和标签,并进行简单的数据拆分:
from sklearn.model_selection import train_test_split
features = ["Pclass", "Sex", "Age", "SibSp", "Parch", "Fare", "Title"]
X = pd.get_dummies(train_df[features]) # 自动处理分类变量
y = train_df["Survived"]
X_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.2, random_state=42)
4.2 模型训练与评估
使用随机森林默认参数进行首次尝试:
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score
model = RandomForestClassifier(random_state=42)
model.fit(X_train, y_train)
predictions = model.predict(X_val)
print(f"验证集准确率: {accuracy_score(y_val, predictions):.4f}")
典型首次运行结果在0.78-0.82之间,说明还有优化空间。
5. 超参数调优进阶
5.1 网格搜索实战
使用GridSearchCV进行参数搜索:
from sklearn.model_selection import GridSearchCV
param_grid = {
'n_estimators': [100, 200, 500],
'max_depth': [None, 5, 10],
'min_samples_split': [2, 5, 10],
'min_samples_leaf': [1, 2, 4]
}
grid_search = GridSearchCV(
estimator=RandomForestClassifier(random_state=42),
param_grid=param_grid,
cv=5,
n_jobs=-1,
verbose=2
)
grid_search.fit(X_train, y_train)
print(f"最佳参数: {grid_search.best_params_}")
print(f"最佳得分: {grid_search.best_score_:.4f}")
提示:n_jobs=-1表示使用所有CPU核心加速计算,大数据集时特别有用。
5.2 特征重要性分析
训练完成后可以查看特征重要性:
import matplotlib.pyplot as plt
feature_importances = pd.DataFrame(
grid_search.best_estimator_.feature_importances_,
index=X.columns,
columns=['importance']
).sort_values('importance', ascending=False)
plt.figure(figsize=(10, 6))
plt.barh(feature_importances.index, feature_importances['importance'])
plt.title("特征重要性排序")
plt.show()
通常会发现性别(Sex)和票价(Fare)是最重要的预测因素。
6. 模型部署与预测
6.1 测试集预测流程
对测试集进行相同的特征工程处理:
# 确保测试集与训练集处理方式完全一致
test_df = fill_na_median(test_df, 'Age')
test_df['Title'] = test_df['Name'].str.extract(' ([A-Za-z]+)\.', expand=False)
test_df['Title'] = test_df['Title'].map(title_mapping)
X_test = pd.get_dummies(test_df[features])
6.2 生成提交文件
predictions = grid_search.best_estimator_.predict(X_test)
output = pd.DataFrame({
'PassengerId': test_df['PassengerId'],
'Survived': predictions
})
output.to_csv('submission.csv', index=False)
print("提交文件已生成")
7. 避坑指南与性能优化
7.1 常见错误排查
-
维度不一致错误 :训练集和测试集get_dummies后列数不同
- 解决方法:先合并两个数据集进行one-hot编码,再拆分
-
过拟合问题 :训练集准确率高但验证集差
- 解决方法:增加min_samples_leaf或使用交叉验证
-
内存不足 :大数据集时网格搜索崩溃
- 解决方法:改用RandomizedSearchCV减少参数组合
7.2 高级优化技巧
- 分类变量编码优化 :尝试Target Encoding代替One-Hot
- 模型融合 :将随机森林与梯度提升树结果加权平均
- 异常值处理 :对Fare列进行对数变换
# Fare列对数变换示例
train_df['Fare'] = train_df['Fare'].map(lambda x: np.log(x) if x > 0 else 0)
8. 项目扩展方向
-
特征工程深度探索 :
- 尝试提取船舱号码中的字母信息
- 利用亲属关系构建家族网络特征
-
模型解释性增强 :
- 使用SHAP值分析单个预测
- 构建决策路径可视化
-
部署为Web服务 :
- 用Flask构建预测API
- 创建交互式预测页面
# SHAP值分析示例代码
import shap
explainer = shap.TreeExplainer(grid_search.best_estimator_)
shap_values = explainer.shap_values(X_val)
shap.summary_plot(shap_values, X_val)
通过这个项目,我深刻体会到特征工程的质量往往比模型选择更重要。在实际业务中,随机森林的稳定表现使其成为我处理结构化数据的首选算法。特别是在数据存在缺失和噪声时,其鲁棒性优势明显。建议初学者不要急于尝试复杂模型,先把基础特征工程和调参技巧练扎实。
更多推荐



所有评论(0)