1. 项目概述与背景解析
泰坦尼克号生存预测项目是机器学习领域的经典入门案例,也是Kaggle平台上最具代表性的竞赛项目之一。这个项目之所以经久不衰,是因为它完美融合了数据科学全流程的各个环节——从数据清洗、特征工程到模型构建与评估,为学习者提供了一个完整的机器学习实战场景。
1912年泰坦尼克号沉船事件中,2224名船员和乘客中有1502人遇难,但生存与否并非完全随机。通过分析乘客的性别、年龄、舱位等级等特征,我们可以构建预测模型来判断特定乘客的生存概率。这个案例的价值在于:
- 数据维度丰富:包含数值型、类别型、文本型等多种数据类型
- 现实意义明确:生存预测结果可以直接验证历史记录
- 问题定义清晰:典型的二分类问题(生存/遇难)
提示:虽然原始数据集来自历史事件,但项目重点在于方法论而非历史分析。所有数据处理和模型构建都应围绕提升预测准确率展开。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据准备与特征工程
2.1 数据集结构解析
原始数据包含两个关键文件:
- train.csv(891条样本):包含特征和标签(Survived)
- test.csv(418条样本):仅包含特征,用于最终预测
关键特征说明:
| 特征名 | 类型 | 描述 | 缺失值情况 |
|---|---|---|---|
| PassengerId | 整型 | 乘客ID | 无 |
| Pclass | 整型 | 舱位等级(1/2/3) | 无 |
| Name | 字符串 | 乘客姓名 | 无 |
| Sex | 字符串 | 性别 | 无 |
| Age | 浮点型 | 年龄 | 约20%缺失 |
| SibSp | 整型 | 同船兄弟姐妹/配偶数 | 无 |
| Parch | 整型 | 同船父母/子女数 | 无 |
| Ticket | 字符串 | 船票编号 | 无 |
| Fare | 浮点型 | 票价 | 少量缺失 |
| Cabin | 字符串 | 舱位编号 | 约77%缺失 |
| Embarked | 字符串 | 登船港口 | 少量缺失 |
2.2 特征筛选与清洗
无效特征剔除:
python复制df.drop(['PassengerId', 'Name', 'Ticket', 'Cabin'], axis=1, inplace=True)
- 乘客ID:纯标识符,无预测价值
- 姓名:虽包含称谓信息但提取复杂,初期可舍弃
- 船票编号:编码规则不统一,难以利用
- 舱位编号:缺失率过高(77%)
缺失值处理策略:
- Age:用中位数填充(受异常值影响小于均值)
python复制df['Age'].fillna(df['Age'].median(), inplace=True)
- Embarked:用众数填充(仅2条缺失)
python复制df['Embarked'].fillna(df['Embarked'].mode()[0], inplace=True)
- Fare:测试集用训练集中对应舱位的平均票价填充
2.3 特征转换与增强
离散化处理:
python复制# 年龄分段
df['AgeGroup'] = pd.cut(df['Age'],
bins=[0, 12, 18, 60, 100],
labels=['Child', 'Teen', 'Adult', 'Elderly'])
# 票价分段
df['FareGroup'] = pd.qcut(df['Fare'], 4,
labels=['Low', 'Medium', 'High', 'Premium'])
类别型特征编码:
python复制# 性别转为0/1
df['Sex'] = df['Sex'].map({'male':0, 'female':1})
# Embarked和新建的分段特征进行one-hot编码
df = pd.get_dummies(df, columns=['Embarked', 'AgeGroup', 'FareGroup'])
特征组合:
python复制# 家庭规模 = 兄弟姐妹数 + 父母子女数 + 1(自己)
df['FamilySize'] = df['SibSp'] + df['Parch'] + 1
# 是否独自旅行
df['IsAlone'] = (df['FamilySize'] == 1).astype(int)
3. 模型构建与训练
3.1 基础模型选择
针对这个二分类问题,我们测试以下经典算法:
| 算法 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 逻辑回归 | 简单、可解释性强 | 线性假设限制 | 基线模型 |
| 随机森林 | 抗过拟合、特征重要性 | 超参敏感 | 结构化数据 |
| GBDT | 预测精度高 | 训练速度慢 | 各类数据 |
| SVM | 小样本效果好 | 大数据量性能差 | 特征维度不高时 |
模型初始化:
python复制from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.svm import SVC
models = {
'LR': LogisticRegression(max_iter=1000),
'RF': RandomForestClassifier(n_estimators=100),
'GBDT': GradientBoostingClassifier(n_estimators=100),
'SVM': SVC(probability=True)
}
3.2 交叉验证策略
采用分层K折交叉验证(StratifiedKFold)确保每折的类别分布一致:
python复制from sklearn.model_selection import StratifiedKFold
skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
评估指标选择:
- 准确率(Accuracy)
- AUC-ROC(综合评估模型排序能力)
- F1-score(平衡精确率与召回率)
3.3 超参数调优示例(以随机森林为例)
使用GridSearchCV进行参数搜索:
python复制param_grid = {
'n_estimators': [50, 100, 200],
'max_depth': [None, 5, 10],
'min_samples_split': [2, 5],
'min_samples_leaf': [1, 2]
}
grid_search = GridSearchCV(
estimator=RandomForestClassifier(random_state=42),
param_grid=param_grid,
cv=skf,
scoring='accuracy',
n_jobs=-1
)
grid_search.fit(X_train, y_train)
注意:实际项目中应先进行粗略搜索确定参数范围,再在小范围内精细调整。过早优化会导致过拟合验证集。
4. 模型评估与解释
4.1 性能对比
经过5折交叉验证后各模型表现:
| 模型 | 平均准确率 | AUC-ROC | F1-score | 训练时间(s) |
|---|---|---|---|---|
| 逻辑回归 | 0.793 | 0.852 | 0.732 | 0.5 |
| 随机森林 | 0.821 | 0.882 | 0.768 | 3.2 |
| GBDT | 0.832 | 0.891 | 0.781 | 6.8 |
| SVM | 0.814 | 0.863 | 0.752 | 12.4 |
4.2 特征重要性分析(随机森林)
python复制importances = best_rf.feature_importances_
indices = np.argsort(importances)[::-1]
plt.figure(figsize=(10,6))
plt.title("Feature Importance")
plt.bar(range(X_train.shape[1]), importances[indices])
plt.xticks(range(X_train.shape[1]), X_train.columns[indices], rotation=90)
plt.show()
典型特征重要性排序:
- Sex(性别)
- Fare(票价)
- Age(年龄)
- Pclass(舱位等级)
- FamilySize(家庭规模)
4.3 模型可解释性技巧
SHAP值分析:
python复制import shap
explainer = shap.TreeExplainer(best_rf)
shap_values = explainer.shap_values(X_test)
shap.summary_plot(shap_values[1], X_test, plot_type="bar")
关键发现:
- 女性(Sex=1)显著提升生存概率
- 高票价(Fare)正向影响生存率
- 三等舱(Pclass=3)降低生存可能
- 儿童(Age<12)有生存优势
5. 系统实现与部署
5.1 预测系统架构设计
code复制[前端界面] -> [Flask API] -> [模型服务]
↓
[结果存储] <- [数据库]
核心组件:
- 前端:HTML/CSS/JavaScript表单
- 后端:Flask处理请求
- 模型:序列化的scikit-learn模型
- 数据库:SQLite存储预测记录
5.2 模型持久化与加载
python复制import joblib
# 保存模型
joblib.dump(best_gbdt, 'titanic_model.pkl')
# 加载模型
model = joblib.load('titanic_model.pkl')
5.3 API接口示例
python复制from flask import Flask, request, jsonify
import pandas as pd
app = Flask(__name__)
@app.route('/predict', methods=['POST'])
def predict():
data = request.json
df = pd.DataFrame([data])
# 执行相同的特征工程
df = preprocess(df)
proba = model.predict_proba(df)[0][1]
return jsonify({'survival_probability': float(proba)})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
5.4 前端表单示例
html复制<form id="predictionForm">
<label>舱位等级:</label>
<select name="Pclass">
<option value="1">一等舱</option>
<option value="2">二等舱</option>
<option value="3">三等舱</option>
</select>
<label>性别:</label>
<select name="Sex">
<option value="0">男性</option>
<option value="1">女性</option>
</select>
<label>年龄:</label>
<input type="number" name="Age" min="0" max="100">
<button type="button" onclick="predict()">预测</button>
</form>
<script>
function predict() {
const formData = Object.fromEntries(new FormData(document.getElementById('predictionForm')));
fetch('/predict', {
method: 'POST',
headers: {'Content-Type': 'application/json'},
body: JSON.stringify(formData)
})
.then(response => response.json())
.then(data => {
alert(`生存概率: ${(data.survival_probability * 100).toFixed(1)}%`);
});
}
</script>
6. 项目进阶方向
6.1 特征工程优化
- 姓名中提取称谓(Mr/Mrs/Miss等)
- 舱位编号的首字母可能包含位置信息
- 船票编号中的字母前缀可能有价值
- 创建家庭ID追踪家庭成员
6.2 模型融合策略
- 投票集成(VotingClassifier)
- 堆叠集成(Stacking)
- 加权平均概率
6.3 自动化机器学习
使用TPOT自动优化管道:
python复制from tpot import TPOTClassifier
tpot = TPOTClassifier(
generations=5,
population_size=20,
cv=skf,
random_state=42,
verbosity=2
)
tpot.fit(X_train, y_train)
6.4 模型监控与更新
- 实现预测结果统计分析面板
- 设置数据漂移检测机制
- 定期用新数据重新训练模型
7. 常见问题与解决方案
7.1 数据相关问题
问题1:年龄缺失值如何处理更合理?
- 进阶方案:用其他特征(如称谓、舱位等)构建回归模型预测缺失年龄
问题2:类别不平衡(约60%遇难)
- 解决方案:
- 过采样少数类(SMOTE)
- 调整类别权重(class_weight='balanced')
- 使用AUC-ROC而非准确率评估
7.2 模型相关问题
问题1:模型在训练集表现好但测试集差
- 检查项:
- 数据泄露(确保测试集未参与任何预处理计算)
- 特征工程一致性(训练/测试应用相同转换)
- 适当增加正则化(如随机森林的max_depth)
问题2:预测概率集中在0.5附近
- 可能原因:
- 模型置信度低
- 特征区分力不足
- 解决方案:
- 尝试更复杂的特征组合
- 使用校准(CalibratedClassifierCV)
7.3 部署相关问题
问题1:API响应慢
- 优化方案:
- 模型轻量化(减小n_estimators)
- 启用批处理预测
- 使用更高效的序列化格式(如ONNX)
问题2:如何保证输入数据质量
- 实现方案:
- 添加输入数据验证
- 对异常值进行自动修正或拒绝
- 记录所有预测请求用于后续分析
8. 项目总结与经验分享
通过这个项目,我深刻体会到特征工程的质量往往比模型选择更重要。几个关键收获:
-
业务理解决定上限:了解泰坦尼克号的历史背景(如"妇女儿童优先")能指导特征创造,如计算家庭成员数量比单独看SibSp/Parch更有意义。
-
迭代比一次完美更重要:我的最佳模型是通过以下迭代过程获得的:
- 第一版:原始特征+逻辑回归(0.76准确率)
- 第二版:基础特征工程+随机森林(0.81)
- 第三版:高级特征组合+GBDT(0.83)
- 最终版:模型集成+超参优化(0.85)
-
可解释性是生产必需:即使GBDT表现最好,在部署时我仍保留了随机森林作为辅助模型,因为其特征重要性更易向非技术人员解释。
-
监控比开发更重要:上线后发现用户常漏填票价字段,导致预测偏差。后来添加了基于舱位的默认值填充逻辑,显著提升了系统稳定性。
对于想进一步挑战的学习者,我建议尝试:
- 将预测系统容器化(Docker)
- 添加基于历史预测的A/B测试功能
- 实现自动特征生成管道(featuretools库)
