泰坦尼克号案例之代码演示
1 课程概览
本课演示泰坦尼克号生存预测案例的完整代码。包括数据加载、数据预处理(缺失值填充)、特征工程、模型训练、模型预测、模型评估和决策树可视化。
2 核心概念与定义
- 数据预处理:缺失值填充、特征提取。
- 特征工程:特征提取、转换。
- 模型训练:使用DecisionTreeClassifier训练。
- 模型评估:使用准确率评估。
- 决策树可视化:绘制决策树图。
3 算法与模型详解
3.1 完整流程
- 加载数据
- 数据预处理
- 提取特征和目标变量
- 缺失值填充
- 特征工程
- 特征转换
- 划分训练集和测试集
- 模型训练
- 模型预测
- 模型评估
- 绘制决策树图
3.2 数据预处理
步骤1:提取特征和目标变量
特征列:
- Pclass:船舱等级
- Sex:性别
- Age:年龄
标签列:
- Survived:存活状态
步骤2:缺失值填充
- Age列有缺失(约200个)
- 使用Age列的平均值填充
注意:
- 不删除缺失值(数据量会减少到700条,太少)
- 使用填充策略
3.3 特征工程
性别转换:
- male → 1
- female → 0
使用OneHot编码或Label编码
3.4 模型训练
API:DecisionTreeClassifier
参数:
- criterion='gini'(CART)
- max_depth:限制深度
- random_state:随机种子
3.5 模型评估
指标:准确率(accuracy_score)
4 代码示例
import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier, export_graphviz
from sklearn.metrics import accuracy_score
from sklearn.feature_extraction import DictVectorizer
# 1. 加载数据
def load_data():
"""加载数据"""
data = pd.read_csv('data/train.csv')
print(f"数据形状: {data.shape}")
print(f"\n前5行:")
print(data.head())
print(f"\n数据信息:")
print(data.info())
print(f"\n缺失值统计:")
print(data.isnull().sum())
return data
# 2. 数据预处理
def preprocess_data(data):
"""数据预处理"""
# 2.1 提取特征和目标变量
X = data[['Pclass', 'Sex', 'Age']]
y = data['Survived']
print("\n=== 提取的特征 ===")
print(X.head())
print("\n=== 标签 ===")
print(y.head())
# 2.2 填充缺失值(Age列用平均值填充)
X.loc[:, 'Age'] = X['Age'].fillna(X['Age'].mean())
print("\n=== 填充后的数据信息 ===")
print(X.info())
return X, y
# 3. 特征工程
def feature_engineering(X):
"""特征工程"""
# 性别转换:male → 1, female → 0
# 方法1:map
# X['Sex'] = X['Sex'].map({'male': 1, 'female': 0})
# 方法2:DictVectorizer(推荐)
# 将DataFrame转为字典,再进行OneHot编码
X_dict = X.to_dict(orient='records')
# 创建DictVectorizer
transfer = DictVectorizer(sparse=False)
X_new = transfer.fit_transform(X_dict)
print("\n=== 特征工程后的数据 ===")
print(f"数据形状: {X_new.shape}")
print(f"特征名称: {transfer.get_feature_names_out()}")
print(f"前5行:\n{X_new[:5]}")
return X_new, transfer
# 4. 模型训练和评估
def train_and_evaluate(X, y):
"""模型训练和评估"""
# 4.1 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.2, random_state=22
)
# 4.2 创建决策树模型
estimator = DecisionTreeClassifier(
criterion='gini',
max_depth=5,
random_state=22
)
# 4.3 模型训练
estimator.fit(X_train, y_train)
# 4.4 模型预测
y_pred = estimator.predict(X_test)
# 4.5 模型评估
accuracy = accuracy_score(y_test, y_pred)
print(f"\n=== 模型评估 ===")
print(f"准确率: {accuracy:.4f}")
return estimator, X_test, y_test
# 5. 决策树可视化
def visualize_tree(estimator, feature_names):
"""决策树可视化"""
# 5.1 导出dot文件
export_graphviz(
estimator,
out_file='tree.dot',
feature_names=feature_names,
class_names=['未存活', '存活'],
filled=True,
rounded=True
)
print("\n决策树已导出到 tree.dot")
# 5.2 转换为图片(需要安装graphviz)
# 命令行执行: dot -Tpng tree.dot -o tree.png
# 5.3 使用matplotlib显示(可选)
try:
from sklearn.tree import plot_tree
import matplotlib.pyplot as plt
plt.figure(figsize=(20, 10))
plot_tree(
estimator,
feature_names=feature_names,
class_names=['未存活', '存活'],
filled=True,
rounded=True
)
plt.title('泰坦尼克号生存预测决策树')
plt.show()
except ImportError:
print("请安装matplotlib: pip install matplotlib")
# 6. 主函数
def main():
"""主函数"""
print("=" * 50)
print("泰坦尼克号生存预测案例")
print("=" * 50)
# 1. 加载数据
data = load_data()
# 2. 数据预处理
X, y = preprocess_data(data)
# 3. 特征工程
X_new, transfer = feature_engineering(X)
# 4. 模型训练和评估
estimator, X_test, y_test = train_and_evaluate(X_new, y)
# 5. 决策树可视化
feature_names = transfer.get_feature_names_out()
visualize_tree(estimator, feature_names)
# 6. 预测新样本
print("\n=== 预测新样本 ===")
# 假设:3等舱,男性,25岁
new_sample = {
'Pclass': 3,
'Sex': 'male',
'Age': 25
}
new_X = transfer.transform([new_sample])
prediction = estimator.predict(new_X)
print(f"3等舱男性25岁: {'存活' if prediction[0] == 1 else '未存活'}")
# 假设:1等舱,女性,30岁
new_sample2 = {
'Pclass': 1,
'Sex': 'female',
'Age': 30
}
new_X2 = transfer.transform([new_sample2])
prediction2 = estimator.predict(new_X2)
print(f"1等舱女性30岁: {'存活' if prediction2[0] == 1 else '未存活'}")
if __name__ == '__main__':
main()
输出示例:
==================================================
泰坦尼克号生存预测案例
==================================================
数据形状: (891, 12)
缺失值统计:
Age 177
Cabin 687
Embarked 2
=== 提取的特征 ===
Pclass Sex Age
0 3 male 22.0
1 1 female 38.0
=== 填充后的数据信息 ===
<class 'pandas.core.frame.DataFrame'>
RangeIndex: 891 entries
Data columns (total 3 columns):
# Column Non-Null Count Dtype
--- ------ -------------- -----
0 Pclass 891 non-null int64
1 Sex 891 non-null object
2 Age 891 non-null float64
=== 特征工程后的数据 ===
数据形状: (891, 4)
特征名称: ['Age' 'Pclass' 'Sex=female' 'Sex=male']
=== 模型评估 ===
准确率: 0.8212
=== 预测新样本 ===
3等舱男性25岁: 未存活
1等舱女性30岁: 存活
5 重难点与易错提醒
- ❗重点:Age列缺失值用平均值填充。
- ❗重点:性别需要转换为数值(OneHot或Label编码)。
- ❗重点:使用DictVectorizer进行特征转换。
- ❗重点:max_depth限制树深度防止过拟合。
- ⚠️易错:fillna的inplace参数在新版本中的警告。
- ⚠️易错:DictVectorizer的sparse参数。
- 💡深入理解:特征工程对模型性能影响很大。
6 课堂问答精选
Q: 如何处理Age列的缺失值?
A: Age列有约200个缺失值,不能删除(数据量会减少到700条,太少)。使用Age列的平均值填充:
X['Age'] = X['Age'].fillna(X['Age'].mean())
Q: 如何将性别转换为数值?
A: 两种方法:
- map方法:
X['Sex'] = X['Sex'].map({'male': 1, 'female': 0}) - DictVectorizer(推荐):将DataFrame转为字典,再进行OneHot编码 推荐使用DictVectorizer,因为它可以自动处理所有类别特征。
7 本课小结
- 数据加载:pd.read_csv
- 数据预处理:缺失值填充(平均值)
- 特征工程:DictVectorizer进行OneHot编码
- 模型训练:DecisionTreeClassifier
- 模型评估:accuracy_score
- 决策树可视化:export_graphviz或plot_tree
8 延伸思考与实践
- 实践:运行泰坦尼克号案例代码。
- 预习:CART决策树回归用法。
- 思考:如何提高模型准确率?