【机器学习-5】 | 泰坦尼克生存预测项目实战篇
0 序言
本文围绕泰坦尼克号生存预测这个项目,详细记录用Python进行数据分析与机器学习的全流程,重点聚焦随机森林算法的项目应用。
内容涵盖数据处理、探索性分析、含逻辑回归、SVM、随机森林多种模型构建及结果评估。读完可掌握从数据到预测的完整流程,理解随机森林在分类任务中的具体实现与应用逻辑。
1 项目背景与目标
1.1 泰坦尼克号事件背景
事件概况:1912年4月15日,泰坦尼克号首航时撞上冰山沉没,2224名乘客和船员中1502人遇难。
悲剧原因:救生艇数量不足是重要原因,且生存存在群体差异,特别是女性、儿童、上层阶级存活率更高。
1.2 目标
通过机器学习工具分析数据,预测哪些乘客幸存,非常经典的一个项目,适合数据分析与机器学习新手实践。
本文展示的内容有:
- Python实现的数据分析全流程,包括
数据处理、可视化、模型构建与评估。 - 重点学习随机森林算法在本项目中的具体应用。
1.3 技术路线概述
为实现预测乘客幸存的目标,我们将按以下逻辑展开:
1.先通过数据处理与探索性分析,明确关键影响因素(如性别、舱位等级等),为模型选择提供依据;
2.从简单模型(逻辑回归)入手,理解分类问题的基础思路;
3.逐步引入更复杂的模型(SVM、随机森林),对比不同算法在本项目中的表现;
4.最终通过评估指标确定最优模型,解释其在泰坦尼克数据中的适用性。
2 数据处理
2.1 所需库
基础库:NumPy(数值计算)、Pandas(数据处理)、IPython(交互式编程)。
机器学习库:SciKit-Learn(模型构建)、SciPy(科学计算)、StatsModels(统计建模)。
可视化与辅助库:Matplotlib(可视化)、Patsy(公式解析)。
import matplotlib.pyplot as plt
%matplotlib inline
import numpy as np
import pandas as pd
import statsmodels.api as sm
from statsmodels.nonparametric.kde import KDEUnivariate
from statsmodels.nonparametric import smoothers_lowess
from pandas import Series, DataFrame
from patsy import dmatrices
from sklearn import datasets, svm
此外,如果要复刻的话,后面还有一个预测的库,可以从链接那获取。
后面做预测的话需要导入该库进行。
2.2 数据导入
import pandas as pd
df = pd.read_csv("data/train.csv") # 读取训练数据
然后你可以直接打印全表看情况:
print(df) # 打印完整数据
当然我不推荐这样子做,上文做分类的时候说过了,
可以只显示前五行基本上就能快速确定数据集里面都有什么。
print(df.head()) # 显示前5行
而且数字可以随便你任意指定,这样子就能防止说数据集太大你全部打印下来,
观感不好!
运行后我们能看到这个数据集的基本情况。


这些数据包含乘客的基本信息,
PassengerId(乘客ID )
Survived(是否幸存,0表示遇难,1表示幸存 )
Pclass(客舱等级 )
Name(姓名 )
Sex(性别 )
Age(年龄 )
SibSp(同乘的兄弟姐妹/配偶数量 )
Parch(同乘的父母/子女数量 )
Ticket(船票编号 )
Fare(船票费用 )
Cabin(客舱号 )
Embarked(登船港口 )
可用于后续如随机森林算法的数据分析与建模,挖掘影响幸存的关键因素 。
若数据显示不全,可按提示点击 scrollable element 滚动查看,或在文本编辑器中打开完整内容 。
这里就不过多展示了,
这一小节的目的就是数据库的初步探索,看看数据是由什么构成的,有什么特征?
方便我们后续操作!
2.3 数据清洗
通过观察发现,数据还是有点复杂的,不像鸢尾花分类那个项目一样,
这里的数据既有缺失值,显示NaN就是,还有一些没有用的特征,
这些都需要进行处理。

处理无用特征:删除缺失值过多的Ticket和Cabin列(对预测价值低)
df = df.drop(['Ticket', 'Cabin'], axis=1) # 删除指定列
处理缺失值:使用dropna()删除含缺失值的行(需先删除无用特征,避免过度丢失数据)
df = df.dropna() # 移除含NaN的行
执行以上程序后,我们得到了一份清洗后的数据集。
通过这些操作减少噪声对后续模型(如逻辑回归、随机森林)的干扰,确保模型学习到真正有价值的规律。
接下来通过探索性分析,挖掘特征与幸存的潜在关系,为模型选择和特征工程提供依据。
3 探索性数据分析
3.1 可视化工具与参数设置
工具:使用Matplotlib,通过plt.figure设置画布大小、分辨率等参数对数据集可视化。
import matplotlib.pyplot as plt
import pandas as pd
from matplotlib.font_manager import FontProperties
# --------------------- 配置中文字体 ---------------------
font = FontProperties(fname=r"C:\Windows\Fonts\msyh.ttc", size=12) # 换成你的字体路径
# 2. 全局设置字体(解决标题、坐标轴、图例的中文显示)
plt.rcParams['font.sans-serif'] = [font.get_name()] # 全局字体
plt.rcParams['axes.unicode_minus'] = False # 解决负号显示为方块的问题
fig = plt.figure(figsize=(18, 6), dpi=1600)
alpha_scatterplot = 0.2
alpha_bar_chart = 0.55
# 子图 1:生存分布
ax1 = plt.subplot2grid((2, 3), (0, 0))
df.Survived.value_counts().plot(kind='bar', alpha=alpha_bar_chart)
ax1.set_xlim(-1, 2)
plt.title("生存分布(1 = 幸存)", fontproperties=font) # 显式指定字体(保险措施)
# 子图 2:年龄与生存
plt.subplot2grid((2, 3), (0, 1))
plt.scatter(df.Survived, df.Age, alpha=alpha_scatterplot)
plt.ylabel("年龄", fontproperties=font)
plt.grid(visible=True, which='major', axis='y')
plt.title("年龄与生存情况分布(1 = 幸存)", fontproperties=font)
# 子图 3:舱位等级分布
ax3 = plt.subplot2grid((2, 3), (0, 2))
df.Pclass.value_counts().plot(kind="barh", alpha=alpha_bar_chart)
ax3.set_ylim(-1, len(df.Pclass.value_counts()))
plt.title("舱位等级分布", fontproperties=font)
# 子图 4:各舱位年龄分布
plt.subplot2grid((2, 3), (1, 0), colspan=2)
df.Age[df.Pclass == 1].plot(kind='kde')
df.Age[df.Pclass == 2].plot(kind='kde')
df.Age[df.Pclass == 3].plot(kind='kde')
plt.xlabel("年龄", fontproperties=font)
plt.title("各舱位乘客年龄分布", fontproperties=font)
plt.legend(('1等舱', '2等舱', '3等舱'), loc='best', prop=font) # 图例字体
# 子图 5:登船港口分布
ax5 = plt.subplot2grid((2, 3), (1, 2))
df.Embarked.value_counts().plot(kind='bar', alpha=alpha_bar_chart)
ax5.set_xlim(-1, len(df.Embarked.value_counts()))
plt.title("各登船港口乘客数量", fontproperties=font)
plt.tight_layout()
plt.show()
运行结果如下:

对该图片进行分析,
图中五个表可以从以下四个方面进行分析:
生存分布:整体存活比例约38%(通过条形图展示)。
性别与生存:女性存活率(约75%)远高于男性(约20%),通过分组条形图验证。
阶级与生存:上层阶级(1-2等舱)存活率高于下层阶级(3等舱),结合性别分析,女性上层阶级存活率最高。
年龄与生存:年龄对生存有影响,通过散点图观察不同年龄段的生存概率变化。
紧接着,我们再单独对我们猜想的特征,一张一张图片生成出来,
再仔细对比一下,看看结果是不是如我们猜想的那样。
plt.figure(figsize=(6,4))
fig, ax = plt.subplots()
df.Survived.value_counts().plot(kind='barh', color="blue", alpha=.65)
ax.set_ylim(-1, len(df.Survived.value_counts()))
plt.title("Survival Breakdown (1 = Survived, 0 = Died)")

这里能看出,泰坦尼克号事故中整体存活率较低!!!
fig = plt.figure(figsize=(18,6))
#create a plot of two subsets, male and female, of the survived variable.
#After we do that we call value_counts() so it can be easily plotted as a bar graph.
#'barh' is just a horizontal bar graph
df_male = df.Survived[df.Sex == 'male'].value_counts().sort_index()
df_female = df.Survived[df.Sex == 'female'].value_counts().sort_index()
ax1 = fig.add_subplot(121)
df_male.plot(kind='barh',label='Male', alpha=0.55)
df_female.plot(kind='barh', color='#FA2379',label='Female', alpha=0.55)
plt.title("Who Survived? with respect to Gender, (raw value counts) "); plt.legend(loc='best')
ax1.set_ylim(-1, 2)
#adjust graph to display the proportions of survival by gender
ax2 = fig.add_subplot(122)
(df_male/float(df_male.sum())).plot(kind='barh',label='Male', alpha=0.55)
(df_female/float(df_female.sum())).plot(kind='barh', color='#FA2379',label='Female', alpha=0.55)
plt.title("Who Survived proportionally? with respect to Gender"); plt.legend(loc='best')
ax2.set_ylim(-1, 2)

这两张图从原始数量看,尽管在原始价值计数中死亡和存活的男性更多,女性幸存人数远多于遇难;
从比例看,女性存活率也显著高于男性。
则共同验证了泰坦尼克号事件中 “性别是影响生存的重要因素”。
接着,我们使用以下程序来进一步对泰坦尼克号数据集(df )进行性别 + 舱位等级 二维交叉分析的可视化,通过 4 个子图,详细拆解不同性别、不同舱位乘客的生存情况。
通过以下程序,我们要先确定好目标:分析“性别 + 舱位等级”组合对生存的影响,
接着去验证3个假设:
- 女性是否普遍比男性存活率高?
- 高舱位(1/2 等舱)是否比低舱位(3 等舱)存活率高?
- 这两个因素是否存在协同作用(比如
女性 + 高舱位存活率极高,男性 + 低舱位存活率极低 )?
fig = plt.figure(figsize=(18,4), dpi=1600)
alpha_level = 0.65
# building on the previous code, here we create an additional subset with in the gender subset
# we created for the survived variable. I know, thats a lot of subsets. After we do that we call
# value_counts() so it it can be easily plotted as a bar graph. this is repeated for each gender
# class pair.
ax1=fig.add_subplot(141)
female_highclass = df.Survived[df.Sex == 'female'][df.Pclass != 3].value_counts()
female_highclass.plot(kind='bar', label='female, highclass', color='#FA2479', alpha=alpha_level)
ax1.set_xticklabels(["Survived", "Died"], rotation=0)
ax1.set_xlim(-1, len(female_highclass))
plt.title("Who Survived? with respect to Gender and Class"); plt.legend(loc='best')
ax2=fig.add_subplot(142, sharey=ax1)
female_lowclass = df.Survived[df.Sex == 'female'][df.Pclass == 3].value_counts()
female_lowclass.plot(kind='bar', label='female, low class', color='pink', alpha=alpha_level)
ax2.set_xticklabels(["Died","Survived"], rotation=0)
ax2.set_xlim(-1, len(female_lowclass))
plt.legend(loc='best')
ax3=fig.add_subplot(143, sharey=ax1)
male_lowclass = df.Survived[df.Sex == 'male'][df.Pclass == 3].value_counts()
male_lowclass.plot(kind='bar', label='male, low class',color='lightblue', alpha=alpha_level)
ax3.set_xticklabels(["Died","Survived"], rotation=0)
ax3.set_xlim(-1, len(male_lowclass))
plt.legend(loc='best')
ax4=fig.add_subplot(144, sharey=ax1)
male_highclass = df.Survived[df.Sex == 'male'][df.Pclass != 3].value_counts()
male_highclass.plot(kind='bar', label='male, highclass', alpha=alpha_level, color='steelblue')
ax4.set_xticklabels(["Died","Survived"], rotation=0)
ax4.set_xlim(-1, len(male_highclass))
plt.legend(loc='best')
在运行结果之前,先对程序进行分析:
这里通过多层布尔索引,筛选出 4 类人群:
female_highclass:女性 + 高舱位(Pclass != 3→ 1/2 等舱 )female_lowclass:女性 + 低舱位(Pclass == 3→ 3 等舱 )male_lowclass:男性 + 低舱位(Pclass == 3)male_highclass:男性 + 高舱位(Pclass != 3)
对每类人群,统计 Survived(0=遇难,1=幸存 )的数量 → value_counts() ,用于后续绘图。
这里程序虽然会生成四个子图,但实际上程序的核心逻辑都是一致的,
这里以子图1的程序为例:
# 1. 创建子图位置(1行4列的第1个位置)
ax1 = fig.add_subplot(141)
# 2. 数据筛选(核心参数:性别=女性,舱位≠3(高舱位))
female_highclass = df.Survived[df.Sex == 'female'][df.Pclass != 3].value_counts()
# 3. 绘图(参数:条形图类型、标签、颜色、透明度)
female_highclass.plot(kind='bar', label='female, highclass', color='#FA2479', alpha=alpha_level)
# 4. 设置X轴标签(参数:标签文本、旋转角度)
ax1.set_xticklabels(["Survived", "Died"], rotation=0)
# 5. 调整X轴范围(参数:左右边界)
ax1.set_xlim(-1, len(female_highclass))
# 6. 设置标题和图例(参数:标题文本、图例位置)
plt.title("Who Survived? with respect to Gender and Class"); plt.legend(loc='best')
同理,后面的子图2、3、4都是一样的道理。
运行以上程序,结果如下:

从左到右,分别对四张子图进行分析:


可能是性别优势抵消了低等级舱位的劣势,再看看下面两种情况。

可能是性别跟舱位的双重劣势。

到这里就很明显了,对比四张图,
我们可以得到以下结论:
无论舱位高低,女性存活率远高于男性,这说明性别是核心影响因素!!!
其次,舱位等级强化性别也有差异,比方说:
- 女性 + 高舱位 → 几乎全部幸存。
- 男性 + 低舱位 → 几乎全部遇难。
结合历史背景验证, 这与泰坦尼克号女士优先、高舱位优先的救援原则完全吻合,说明数据背后的逻辑符合真实事件。
现在我们有了更多关于谁在这场悲剧中幸存和死亡的信息。有了这种更深入的理解,我们能更好地创建更优秀的模型。
而这也是交互式数据分析中的典型过程!!!
先从小处着手,了解最基本的关系,随着发现越来越多的数据信息,再慢慢增加分析的复杂性。
具体可以参考以下程序:
fig = plt.figure(figsize=(18,12), dpi=1600)
a = 0.65
# Step 1
ax1 = fig.add_subplot(341)
df.Survived.value_counts().plot(kind='bar', color="blue", alpha=a)
ax1.set_xlim(-1, len(df.Survived.value_counts()))
plt.title("Step. 1")
# Step 2
ax2 = fig.add_subplot(345)
df.Survived[df.Sex == 'male'].value_counts().plot(kind='bar',label='Male')
df.Survived[df.Sex == 'female'].value_counts().plot(kind='bar', color='#FA2379',label='Female')
ax2.set_xlim(-1, 2)
plt.title("Step. 2 \nWho Survived? with respect to Gender."); plt.legend(loc='best')
ax3 = fig.add_subplot(346)
(df.Survived[df.Sex == 'male'].value_counts()/float(df.Sex[df.Sex == 'male'].size)).plot(kind='bar',label='Male')
(df.Survived[df.Sex == 'female'].value_counts()/float(df.Sex[df.Sex == 'female'].size)).plot(kind='bar', color='#FA2379',label='Female')
ax3.set_xlim(-1,2)
plt.title("Who Survied proportionally?"); plt.legend(loc='best')
# Step 3
ax4 = fig.add_subplot(349)
female_highclass = df.Survived[df.Sex == 'female'][df.Pclass != 3].value_counts()
female_highclass.plot(kind='bar', label='female highclass', color='#FA2479', alpha=a)
ax4.set_xticklabels(["Survived", "Died"], rotation=0)
ax4.set_xlim(-1, len(female_highclass))
plt.title("Who Survived? with respect to Gender and Class"); plt.legend(loc='best')
ax5 = fig.add_subplot(3,4,10, sharey=ax1)
female_lowclass = df.Survived[df.Sex == 'female'][df.Pclass == 3].value_counts()
female_lowclass.plot(kind='bar', label='female, low class', color='pink', alpha=a)
ax5.set_xticklabels(["Died","Survived"], rotation=0)
ax5.set_xlim(-1, len(female_lowclass))
plt.legend(loc='best')
ax6 = fig.add_subplot(3,4,11, sharey=ax1)
male_lowclass = df.Survived[df.Sex == 'male'][df.Pclass == 3].value_counts()
male_lowclass.plot(kind='bar', label='male, low class',color='lightblue', alpha=a)
ax6.set_xticklabels(["Died","Survived"], rotation=0)
ax6.set_xlim(-1, len(male_lowclass))
plt.legend(loc='best')
ax7 = fig.add_subplot(3,4,12, sharey=ax1)
male_highclass = df.Survived[df.Sex == 'male'][df.Pclass != 3].value_counts()
male_highclass.plot(kind='bar', label='male highclass', alpha=a, color='steelblue')
ax7.set_xticklabels(["Died","Survived"], rotation=0)
ax7.set_xlim(-1, len(male_highclass))
plt.legend(loc='best')
以上程序生成一个泰坦尼克号生存分析的“分步可视化报告”,一共有7个子图。
通过这 7 个子图,从简单到复杂、从单特征到多特征交叉,逐步拆解生存规律。
主要逻辑是:先看整体生存分布→再拆性别影响→最后看性别 + 舱位的协同影响。
| 步骤(Step) | 分析维度 | 子图数量 | 核心目标 |
|---|---|---|---|
| Step 1 | 整体生存分布 | 1 张 | 看全局:幸存 vs 遇难的总人数 |
| Step 2 | 性别对生存的影响 | 2 张 | 拆差异:男女生存数量+比例 |
| Step 3 | 性别+舱位对生存的影响 | 4 张 | 挖细节:不同性别+舱位的生存 |
而为什么会想到这个步骤,自然也是通过前面可视化图的对比得出来的结论。
下面针对该程序,对程序里的内容进行分段分析。
1.整体框架与画布布局
fig = plt.figure(figsize=(18,12), dpi=1600)
a = 0.65
plt.figure(...)创建一个画布,指定尺寸为 18x12 英寸,分辨率 1600(数值越大图越清晰,也更占内存 )。a 定义了后续绘图的透明度,让图表不会因颜色过深显得厚重。
2.整体生存分布可视化(最宏观)
ax1 = fig.add_subplot(341)
df.Survived.value_counts().plot(kind='bar', color="blue", alpha=a)
plt.title("Step. 1")
fig.add_subplot(341):在3行4列布局的第 1 个位置(从左到右、从上到下数 ),创建一个子图对象ax1,后续绘图都基于它。df.Survived.value_counts():先统计Survived列(是否幸存,值为0或1)里,0和1各自出现的次数。.plot(kind='bar'...):把统计结果绘制成垂直条形图,蓝色填充,透明度用前面定义的a(0.65),这样图看起来不会太实,有层次感。plt.title("Step. 1"):给这个子图加标题,标识这是分析的第一步。- 运行后会输出一张展示
Survived字段取值分布的条形图,能直观看到遇难和幸存对应的人数高低。
3.性别对生存的影响可视化
子图 2 - 男女生存数量对比(绝对数)
ax2 = fig.add_subplot(345)
df.Survived[df.Sex == 'male'].value_counts().plot(kind='bar',label='Male')
df.Survived[df.Sex == 'female'].value_counts().plot(kind='bar', color='#FA2379',label='Female')
plt.title("Step. 2 \nWho Survived? with respect to Gender.");
- 程序逻辑:
fig.add_subplot(345):在3行4列布局的第 5 个位置(第 2 行第 1 列 )创建子图ax2。df.Survived[df.Sex == 'male'].value_counts():先筛选出Sex为male的行,再统计这些行里Survived列0和1的数量,也就是男性群体中遇难和幸存的人数。同理,后面是统计女性群体的。.plot(kind='bar'...):分别把男性、女性的统计结果绘制成垂直条形图,男性用默认颜色(一般是蓝色系 ),女性用#FA2379(偏粉色 )区分,同时给各自加上图例标签Male、Female,方便看图标注。plt.title(...):给子图加标题,说明这一步是分析性别对生存的影响。
- 运行后会输出一张有两组条形(分别代表男性、女性 )的图,每组又有两根条(对应
0遇难、1幸存 ),能直观对比男性和女性群体里,遇难、幸存人数的绝对数量差异。
子图 3 - 男女生存比例对比(相对数)
ax3 = fig.add_subplot(346)
(df.Survived[df.Sex == 'male'].value_counts()/float(df.Sex[df.Sex == 'male'].size)).plot(kind='bar',label='Male')
(df.Survived[df.Sex == 'female'].value_counts()/float(df.Sex[df.Sex == 'female'].size)).plot(kind='bar', color='#FA2379',label='Female')
plt.title("Who Survied proportionally?");
- 程序逻辑:
fig.add_subplot(346):在3行4列布局的第 6 个位置(第 2 行第 2 列 )创建子图ax3。(df.Survived[df.Sex == 'male'].value_counts()/float(df.Sex[df.Sex == 'male'].size)):分子是男性群体中遇难、幸存的人数,分母是男性群体的总人数(df.Sex[df.Sex == 'male'].size统计男性行数 ),相除得到男性群体中遇难、幸存的比例。同理计算女性群体的比例。.plot(kind='bar'...):把比例结果绘制成垂直条形图,颜色、图例标签和前面子图 2 对应,方便对比。plt.title(...):加标题说明这是比例对比。
和子图 2 布局类似,但条形代表的是比例。这样能消除男女人数基数不同的影响,更公平地对比男女生存率。比如若女性总人数少,但幸存比例高,就能清晰展现出来。
4.性别 + 舱位对生存的影响(更微观角度)
# 女高舱(1/2等舱)
ax4 = fig.add_subplot(349)
female_highclass = df.Survived[df.Sex == 'female'][df.Pclass != 3].value_counts()
female_highclass.plot(kind='bar', label='female highclass', color='#FA2479', alpha=a)
ax4.set_xticklabels(["Survived", "Died"], rotation=0)
ax4.set_xlim(-1, len(female_highclass))
plt.title("Who Survived? with respect to Gender and Class"); plt.legend(loc='best')
# 女低舱(3等舱)
ax5 = fig.add_subplot(3,4,10)
female_lowclass = df.Survived[df.Sex == 'female'][df.Pclass == 3].value_counts()
female_lowclass.plot(kind='bar', label='female, low class', color='pink', alpha=a)
ax5.set_xticklabels(["Died","Survived"], rotation=0)
ax5.set_xlim(-1, len(female_lowclass))
plt.legend(loc='best')
# 男低舱(3等舱)
ax6 = fig.add_subplot(3,4,11)
male_lowclass = df.Survived[df.Sex == 'male'][df.Pclass == 3].value_counts()
male_lowclass.plot(kind='bar', label='male, low class',color='lightblue', alpha=a)
ax6.set_xticklabels(["Died","Survived"], rotation=0)
ax6.set_xlim(-1, len(male_lowclass))
plt.legend(loc='best')
# 男高舱(1/2等舱)
ax7 = fig.add_subplot(3,4,12)
male_highclass = df.Survived[df.Sex == 'male'][df.Pclass != 3].value_counts()
male_highclass.plot(kind='bar', label='male highclass', alpha=a, color='steelblue')
ax7.set_xticklabels(["Died","Survived"], rotation=0)
ax7.set_xlim(-1, len(male_highclass))
plt.legend(loc='best')
- 这里以女高舱
ax4为例,其他类似:fig.add_subplot(349):在3行4列布局的第 9 个位置(第 3 行第 1 列 )创建子图ax4。df.Survived[df.Sex == 'female'][df.Pclass != 3]:先筛选出女性且**舱位不是 3 等舱(即 1/2 等舱,视为高舱 )**的数据行,再统计这些行里Survived列遇难和幸存的数量。.plot(kind='bar'...):用特定颜色绘制成垂直条形图,ax4.set_xticklabels(["Survived", "Died"], rotation=0):把 X 轴默认的0、1刻度标签,替换成更易懂的Survived(幸存 )、Died(遇难 ),旋转角度设为0(即水平显示 ),方便阅读。ax4.set_xlim(-1, len(female_highclass)):调整 X 轴范围,让条形和边界有一定空隙,图更美观。plt.legend(loc='best'):自动把图例放到图里最适合的位置。
- 运行后会输出4 张子图,分别对应
女性高舱\女性低舱\男性低舱\男性高舱群体的生存情况。我们能从图中看到不同性别和舱位组合下,遇难、幸存人数的差异。
接下来看一下运行结果:

3.2 探索性分析结论与模型选择启示
通过可视化分析,我们发现:
性别、舱位等级是强影响因素(女性、高舱位存活率显著更高);
年龄、登船港口等特征与生存存在一定关联,但关系非线性(如不同年龄段的生存概率波动较大)。
这些发现对模型选择的启示有以下三点:
1.逻辑回归适合捕捉线性关系(如 “女性→高存活率” 的直接关联),可作为基础模型验证核心特征的影响;
2.SVM 能处理非线性关系,但对特征缩放敏感,需结合数据预处理;
3.随机森林擅长捕捉特征间的复杂交互,且对缺失值和非线性关系的适应性更强,适合本项目的多特征场景。
综上所述,我们可以得出最后的结论:性别是最关键因素,舱位等级是次要但强化的因素!!!
经过可视化分析,我们清晰看到性别、舱位等特征与幸存强相关。
但可视化只能描述规律,无法回答核心问题:给定一位乘客的特征(如女性、高舱位),如何精准预测她是否会幸存?
这属于二元分类任务(0/1 预测),而逻辑回归,正是解决这类问题的经典方法 —— 它通过概率建模,能从历史数据中学习特征如何影响幸存概率,并自动找到区分 0/1 的最优阈值,让预测从个人的主观判断 变为计算机的数据驱动。
接下来,我们就用逻辑回归,构建第一个预测模型,看看它如何从数据中学习规律,也方便引出后续随机森林这类进阶模型。
4 机器学习模型构建
基于探索性分析发现的线性关联,如性别、舱位与生存的直接关系,我们先使用逻辑回归模型,它能直观展示特征对生存的影响权重,是分类问题的基础工具。
4.1 逻辑回归
原理:用于预测二分类结果,通过逻辑函数计算生存概率,设定阈值转化为分类结果。
实现步骤:
构建公式:Survived ~ C(Pclass) + C(Sex) + Age + SibSp + C(Embarked)数据转换:用dmatrices处理分类变量为布尔值模型训练与评估:
#定义逻辑回归的特征公式,明确用哪些特征(Pclass/Sex/Age 等 )预测 Survived。
#准备 results 字典,后续把模型训练结果存进去。
formula = 'Survived ~ C(Pclass) + C(Sex) + Age + SibSp + C(Embarked)'
results = {}
# 1. 用 patsy 的 dmatrices 处理数据
# 根据 formula 把原始数据(df)转成模型友好的格式:
# - y:预测目标(Survived)
# - x:处理后的特征矩阵(分类变量会转成虚拟变量,如 Sex 转成 Sex_male/Sex_female )
y, x = dmatrices(formula, data=df, return_type='dataframe')
# 2. 实例化逻辑回归模型
# sm.logit 是 statsmodels 里的逻辑回归实现(基于对数几率回归)
model = sm.logit(y, x)
# 3. 训练模型(拟合数据)
res = model.fit()
# 4. 把训练结果存入 results 字典
# 存了模型结果(res)和公式(formula),方便后续调用(如预测、对比)
results['Logit'] = [res, formula]
# 5. 输出模型摘要(关键!)
# 会打印系数、P 值、模型指标(如 AIC、BIC ),用于分析模型效果
res.summary()
以上程序是用 patsy 处理数据成模型可识别的格式,再用 statsmodels 拟合逻辑回归,最后输出模型结果。
结果如下:

这里主要看上图中红圈框出来的内容,
其他的如:
- 截距项(Intercept)
coef=4.5423,P>|z|=0.000(显著 )
当所有特征取基准值时,模型预测的幸存对数几率为 4.54。 - 舱位等级(C (Pclass)[T.2]、C (Pclass)[T.3])
C(Pclass)[T.2]:coef=-1.2673,P>|z|=0.000
相对于1 等舱,2 等舱乘客的幸存对数几率降低 1.27 → 2 等舱幸存概率更低。
C(Pclass)[T.3]:coef=-2.4966,P>|z|=0.000
相对于1 等舱,3 等舱乘客的幸存对数几率降低 2.50 → 3 等舱幸存概率远低于 1 等舱,影响比 2 等舱更大。 - 性别(C (Sex)[T.male])
coef=-2.6239,P>|z|=0.000
相对于女性,男性乘客的幸存对数几率降低 2.62 → 男性幸存概率远低于女性。 - 登船港口(C (Embarked)[T.Q]、C (Embarked)[T.S])
C(Embarked)[T.Q]:coef=-0.8351,P>|z|=0.162(不显著 )
相对于南安普顿(S,基准),从皇后镇(Q)登船的乘客,“幸存对数几率” 降低 0.84,但因 P>|z|>0.05,港口 Q 的影响无统计学意义(可能样本量少或与其他特征重叠 )。
C(Embarked)[T.S]:coef=-0.4254,P>|z|=0.116(不显著 )
相对于南安普顿(S,基准),从瑟堡登船的乘客,“幸存对数几率” 降低 0.43,但无统计学意义。 - 年龄(Age)
coef=-0.0436,P>|z|=0.000(显著 )
年龄每增加 1 岁,幸存对数几率降低 0.044 → 年龄越大,幸存概率越低。 - 亲属同行数量(SibSp)
coef=-0.3697,P>|z|=0.003(显著 )
同行亲属 / 配偶数量每增加 1 人,“幸存对数几率” 降低 0.37 → 亲属越多,幸存概率越低,可能因救援时需照顾家人,或低舱位乘客亲属更多。
接着我们计算逻辑回归模型的预测结果和残差,帮助评估模型拟合效果。
plt.figure(figsize=(18,4))
plt.subplot(121, facecolor="#5E5656")
ypred = res.predict(x)
plt.plot(x.index, ypred, 'bo', x.index, y, 'mo', alpha=.25)
plt.grid(color='white', linestyle='dashed')
plt.title('Logit predictions, Blue: \nFitted/predicted values: Red')
ax2 = plt.subplot(122, facecolor="#5E5656")
plt.plot(res.resid_dev, 'r-')
plt.grid(color='white', linestyle='dashed')
ax2.set_xlim(-1, len(res.resid_dev))
plt.title('Logit Residuals')
运行结果如下:

从图像可以看出,模型对极端概率样本预测较准,残差无明显趋势,模型结构基本合理。
但是对中间概率样本预测能力不足,且存在部分预测误差大的样本,需进一步优化特征或模型。
再用以下程序来辅助理解特征与幸存概率的关系 以及 预测值的分布规律。
fig = plt.figure(figsize=(18,9), dpi=1600)
a = .2
# Below are examples of more advanced plotting.
# It it looks strange check out the tutorial above.
fig.add_subplot(221, facecolor ="#DBDBDB")
kde_res = KDEUnivariate(res.predict())
kde_res.fit()
plt.plot(kde_res.support,kde_res.density)
plt.fill_between(kde_res.support,kde_res.density, alpha=a)
plt.title("Distribution of our Predictions")
fig.add_subplot(222, facecolor ="#DBDBDB")
plt.scatter(res.predict(),x['C(Sex)[T.male]'] , alpha=a)
plt.grid(visible=True, which='major', axis='x')
plt.xlabel("Predicted chance of survival")
plt.ylabel("Gender Bool")
plt.title("The Change of Survival Probability by Gender (1 = Male)")
fig.add_subplot(223, facecolor ="#DBDBDB")
plt.scatter(res.predict(),x['C(Pclass)[T.3]'] , alpha=a)
plt.xlabel("Predicted chance of survival")
plt.ylabel("Class Bool")
plt.grid(visible=True, which='major', axis='x')
plt.title("The Change of Survival Probability by Lower Class (1 = 3rd Class)")
fig.add_subplot(224, facecolor ="#DBDBDB")
plt.scatter(res.predict(),x.Age , alpha=a)
plt.grid(True, linewidth=0.15)
plt.title("The Change of Survival Probability by Age")
plt.xlabel("Predicted chance of survival")
plt.ylabel("Age")
结果如下:

整体来说:模型拟合符合预期
性别、舱位、年龄对幸存概率的影响,均与逻辑回归的系数解读一致,说明模型学到了业务常识。
预测值的双峰分布,反映模型对多数样本能明确判断,但中间区间的模糊样本仍存在。
接下来我们使用我们训练好的模型去预测数据,看看预测出来的数据效果如何。
test_data = pd.read_csv("data/test.csv")
print(test_data)

由于部分特征(如年龄)与生存的关系是非线性的,我们引入 SVM 模型。
它通过核函数将数据映射到高维空间,可捕捉更复杂的边界关系,弥补逻辑回归的线性局限。
4.2 支持向量机(SVM)
原理:通过核函数将数据映射到高维空间,寻找最优分类超平面,支持线性、RBF、多项式等核函数。
实现步骤:
- 数据准备:提取特征矩阵
X和目标向量y,打乱并划分训练/测试集 - 多核函数对比:
# set plotting parameters
plt.figure(figsize=(8,6))
# create a regression friendly data frame
y, x = dmatrices(formula_ml, data=df, return_type='matrix')
# select which features we would like to analyze
# try chaning the selection here for diffrent output.
# Choose : [2,3] - pretty sweet DBs [3,1] --standard DBs [7,3] -very cool DBs,
# [3,6] -- very long complex dbs, could take over an hour to calculate!
feature_1 = 2
feature_2 = 3
X = np.asarray(x)
X = X[:,[feature_1, feature_2]]
y = np.asarray(y)
# needs to be 1 dimenstional so we flatten. it comes out of dmatirces with a shape.
y = y.flatten()
n_sample = len(X)
np.random.seed(0)
order = np.random.permutation(n_sample)
X = X[order]
y = y[order].astype(float)
# do a cross validation
nighty_precent_of_sample = int(.9 * n_sample)
X_train = X[:nighty_precent_of_sample]
y_train = y[:nighty_precent_of_sample]
X_test = X[nighty_precent_of_sample:]
y_test = y[nighty_precent_of_sample:]
# create a list of the types of kerneks we will use for your analysis
types_of_kernels = ['linear', 'rbf', 'poly']
# specify our color map for plotting the results
color_map = plt.cm.RdBu_r
# fit the model
for fig_num, kernel in enumerate(types_of_kernels):
clf = svm.SVC(kernel=kernel, gamma=3)
clf.fit(X_train, y_train)
plt.figure(fig_num)
plt.scatter(X[:, 0], X[:, 1], c=y, zorder=10, cmap=color_map)
# circle out the test data
plt.scatter(X_test[:, 0], X_test[:, 1], s=80, facecolors='none', zorder=10)
plt.axis('tight')
x_min = X[:, 0].min()
x_max = X[:, 0].max()
y_min = X[:, 1].min()
y_max = X[:, 1].max()
XX, YY = np.mgrid[x_min:x_max:200j, y_min:y_max:200j]
Z = clf.decision_function(np.c_[XX.ravel(), YY.ravel()])
# put the result into a color plot
Z = Z.reshape(XX.shape)
plt.pcolormesh(XX, YY, Z > 0, cmap=color_map)
plt.contour(XX, YY, Z, colors=['k', 'k', 'k'], linestyles=['--', '-', '--'],
levels=[-.5, 0, .5])
plt.title(kernel)
plt.show()
特点:非线性核函数(如RBF)可捕捉更复杂的数据结构,但解释性较线性模型差。
该程序的核心是 用支持向量机(SVM)对二维特征进行二分类,并可视化不同核函数的决策边界,帮助理解 SVM 的核函数对分类结果的影响。
1. 数据准备:生成设计矩阵
y, x = dmatrices(formula_ml, data=df, return_type='matrix')
- 工具:
patsy.dmatrices(用于将公式和数据转换为模型友好的矩阵) - 作用:
- 自动处理 分类变量(如
C(Pclass)会生成虚拟变量)、截距项(默认添加第一列全1)。 - 返回
y(标签,二维列向量)和x(特征矩阵,含截距)。
- 自动处理 分类变量(如
2. 特征选择:降维到二维
feature_1 = 2
feature_2 = 3
X = np.asarray(x)[:, [feature_1, feature_2]]
y = np.asarray(y).flatten()
- 目的:仅选择 两个特征(方便在二维平面可视化决策边界)。
- 细节:
x的第一列是截距(全1),所以feature_1=2对应 第三个特征(索引从0开始)。y.flatten()将列向量转为一维数组(SVM 要求标签是一维)。
3. 数据划分:随机打乱+训练测试拆分
np.random.seed(0)
order = np.random.permutation(n_sample) # 随机打乱样本索引
X = X[order]
y = y[order].astype(float)
nighty_precent_of_sample = int(.9 * n_sample) # 90%训练集,10%测试集
X_train, y_train = X[:nighty_precent_of_sample], y[:nighty_precent_of_sample]
X_test, y_test = X[nighty_precent_of_sample:], y[nighty_precent_of_sample:]
- 关键:
np.random.seed(0)保证结果可复现。
4. 模型训练+可视化:对比3种核函数
types_of_kernels = ['linear', 'rbf', 'poly'] # 三种核函数
color_map = plt.cm.RdBu_r # 红蓝配色,区分两类
for fig_num, kernel in enumerate(types_of_kernels):
clf = svm.SVC(kernel=kernel, gamma=3) # gamma控制核函数复杂度(越大越复杂)
clf.fit(X_train, y_train) # 训练模型
# 可视化步骤:
plt.figure(fig_num)
# 1. 绘制所有数据点(颜色区分类别)
plt.scatter(X[:, 0], X[:, 1], c=y, zorder=10, cmap=color_map)
# 2. 突出测试集(空心圆)
plt.scatter(X_test[:, 0], X_test[:, 1], s=80, facecolors='none', zorder=10)
# 3. 计算决策边界的网格
x_min, x_max = X[:, 0].min(), X[:, 0].max()
y_min, y_max = X[:, 1].min(), X[:, 1].max()
XX, YY = np.mgrid[x_min:x_max:200j, y_min:y_max:200j] # 200x200网格
Z = clf.decision_function(np.c_[XX.ravel(), YY.ravel()]) # 计算每个网格点的决策值
# 4. 绘制决策区域(颜色)和决策边界( contour )
Z = Z.reshape(XX.shape)
plt.pcolormesh(XX, YY, Z > 0, cmap=color_map) # Z>0 为一类,否则另一类
plt.contour(XX, YY, Z, colors=['k', 'k', 'k'], linestyles=['--', '-', '--'], levels=[-.5, 0, .5])
plt.title(kernel)
plt.show()
运行结果如下:



5.核函数对比
| 核函数 | 决策边界形状 | 特点 | 代码中表现(gamma=3) |
|---|---|---|---|
| linear | 直线 | 简单,计算快,适合线性可分数据 | 边界平直,对复杂分布拟合能力弱 |
| rbf | 曲线(径向对称) | 非线性,适合复杂分布,gamma影响大 | 边界弯曲,贴合数据(gamma=3时较复杂) |
| poly | 多项式曲线 | 非线性,degree默认=3(代码未显式设,需注意) | 边界可能更“扭曲”,拟合更细 |
考虑到多个特征(如性别 + 舱位 + 年龄)可能存在交互作用(例如女性在高舱位的存活率远高于单一特征的叠加),我们最终使用随机森林模型。
它通过多棵决策树的集成,能自动学习特征间的协同效应,且抗过拟合能力更强,适合本项目的复杂场景。
4.3 随机森林(核心)
原理:集成学习方法,通过构建多个决策树,取多数结果作为最终预测(分类任务),利用“群体智慧”降低过拟合风险。
实现步骤:
- 导入库与数据准备:
import sklearn.ensemble as ske
y, x = dmatrices(formula_ml, data=df, return_type='matrix')
y = y.flatten() # 转换为一维数组
- 模型训练与评分:
rf = ske.RandomForestClassifier()
rf.fit(x, y)
score = rf.score(x, y) # 计算准确率,示例结果为0.945

这里可以看到,随机森林准确率是要高于逻辑回归的。
5 模型评估与对比
为判断不同模型的适用性,我们使用准确率、精确率、召回率及混淆矩阵进行评估:
逻辑回归:在捕捉线性关系上表现稳定,但对非线性特征的利用不足;
SVM:非线性拟合能力优于逻辑回归,但对参数敏感,计算成本较高;
随机森林:在本项目中表现最优,尤其在处理多特征交互和噪声数据上优势明显,这与探索性分析中发现的`多因素协同影响生存``结论一致。
因此,随机森林更适合作为本项目的最终预测模型。
6 小结
本笔记围绕泰坦尼克号生存预测项目,从数据处理到模型构建展开,重点解析了随机森林算法的应用。流程上,先通过清洗与探索性分析理解数据特征,再依次实现逻辑回归、SVM、随机森林模型,最终通过交叉验证和结果输出。
随机森林作为集成方法,在本项目中表现出较高准确率,其核心在于通过多棵决策树的集成降低误差,适合处理复杂特征关系。通过学习,可掌握从实际问题到机器学习解决方案的完整流程,理解不同模型的适用场景与实现细节。
关于前文提到的predict.py文件,我有对其进行修改,所以还是给个获取链接。
以下是获取链接:predict文件链接-百度网盘
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)