刚开始接触机器学习的时候,大多数人第一个上手的项目就是鸢尾花分类。这个数据集在 sklearn 里内置,一行代码就能加载,150 条样本、4 个特征、3 个类别,结构简单但信息量足够丰富,特别适合拿来熟悉数据分析和可视化的基本操作。但有一个问题很常见:数据加载出来了,看到的是一堆数字矩阵,怎么把三类样本在图上区分开?这篇文章我整理了五种我实际用过的方案,从最基础的 matplotlib 手写散点图,到一行代码出完整图的 seaborn pairplot,再到可交互的 3D 图,每套代码都是可以直接复制的,你按需取用就行。
1. 准备工作:先认识鸢尾花数据集
1.1 鸢尾花数据集为什么适合入门
鸢尾花数据集(Iris Dataset)几乎可以说是机器学习界的“Hello World”。它由英国统计学家 Ronald Fisher 在 1936 年引入,包含 3 个鸢尾花品种(setosa、versicolor、virginica)各 50 条样本,每条样本记录 4 个特征:
- sepal length (cm):花萼长度
- sepal width (cm):花萼宽度
- petal length (cm):花瓣长度
- petal width (cm):花瓣宽度
这个数据集的巧妙之处在于:setosa 这个品种和另外两个品种在花瓣特征上区分度极高,而 versicolor 和 virginica 之间有部分重叠,既不是完全线性可分,也不是完全揉成一团。用这个数据集来做可视化练习,你能非常直观地感受到“特征选择”的重要性——选对了特征组合,类别边界一目了然;选错了,三类样本混在一起怎么看都分不开。
另外,sklearn 内置了这个数据集,不需要去网上下载文件,也不需要联网,本地环境装好 sklearn 就能直接用。对于刚入门、还在折腾环境的新手来说,省去了数据获取这一步,可以把精力集中在理解数据和处理流程上。
1.2 用一行代码加载数据集
sklearn 提供了 load_iris() 函数,使用方式非常简单:
python复制from sklearn.datasets import load_iris
iris = load_iris()
print(type(iris))
# <class 'sklearn.utils.Bunch'>
这里返回的是一个 Bunch 对象,你可以把它理解成一个“加强版字典”,里面同时包含了数据、标签、特征名、类别名等多个部分。直接用 dir(iris) 或者 .keys() 就能看到它有哪些字段。
python复制print(iris.keys())
# dict_keys(['data', 'target', 'frame', 'target_names', 'DESCR', 'feature_names', 'filename', 'data_module'])
我常用的字段有几个:data 是特征矩阵,形状为 (150, 4);target 是类别标签的数值数组,0、1、2 分别对应三个品种;target_names 是类别名的字符串数组;feature_names 是四个特征名称的列表。这几个字段配合使用,基本可以完成大部分数据探索和可视化工作。
1.3 转换成 DataFrame 并确认数据结构
虽然直接用 numpy 数组也能做可视化,但我还是建议先把数据转成 pandas 的 DataFrame,这样查看数据、筛选样本、和 seaborn 等可视化库配合起来都更顺手。
python复制import pandas as pd
df = pd.DataFrame(iris.data, columns=iris.feature_names)
df['label'] = iris.target
df['species'] = iris.target_names[iris.target]
print(df.head())
运行这段代码,你会看到前 5 行数据,每一行代表一朵鸢尾花。label 列是数值标签,species 列是具体的品种名称。我通常在 DataFrame 里同时保留这两列:label 用于按数值筛选或作为颜色映射,species 用于图例展示和分组操作,两边都方便。
再确认一下类别分布:
python复制print(df['species'].value_counts())
# setosa 50
# versicolor 50
# virginica 50
三个类别各 50 条,非常均衡,这就不用担心可视化时某一类样本过少导致图形失真了。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 方案一:matplotlib 手写散点图,最基础也最可控
2.1 为什么要先看手写方案
现在有很多现成的封装好的画图函数,两三行代码就能出一张漂亮的图。但如果你想真正理解散点图是怎么画出来的、每个点是怎么根据特征值落到坐标轴上的,那还是得亲手写一遍。方案一就是用 matplotlib 的 scatter 函数,通过循环遍历三个类别,分别画三种颜色的点。这个方案看起来啰嗦,但它是所有后续方案的基础,也是排查问题时最好定位的一个方案。
2.2 核心代码实现
python复制import matplotlib.pyplot as plt
colors = ['#FF6B6B', '#4ECDC4', '#45B7D1']
# 选择两个特征进行可视化,这里选的是花萼长度和花瓣长度
x_idx, y_idx = 0, 2
x_label, y_label = iris.feature_names[x_idx], iris.feature_names[y_idx]
plt.figure(figsize=(8, 6))
for i, label in enumerate(iris.target_names):
# 布尔索引:选出当前类别的样本
mask = iris.target == i
plt.scatter(
iris.data[mask, x_idx],
iris.data[mask, y_idx],
c=colors[i],
label=label,
alpha=0.8,
edgecolors='white',
linewidths=0.5,
s=50,
)
plt.xlabel(x_label)
plt.ylabel(y_label)
plt.title(f'Iris Dataset: {x_label} vs {y_label}')
plt.legend()
plt.grid(alpha=0.3)
plt.show()
2.3 这段代码里的几个关键细节
mask = iris.target == i 这行是筛选样本的核心思路。它的原理是生成一个布尔数组,长度和样本数一致,True 表示当前样本属于类别 i。然后用 iris.data[mask, x_idx] 就可以取出该类别在指定特征维度上的所有值。这个布尔索引写法在数据处理中非常常用,建议新手把它记牢。
alpha=0.8 控制点的透明度,当两个类的点重叠较多时,透明度能帮你看出哪些区域分布密集。edgecolors='white' 给每个点加上白色描边,点多了以后不容易糊成一片,视觉上更干净。s=50 是点的大小,如果在数据量更大的场景里,可以适当调小一点。
至于为什么我选了第 0 个特征(sepal length)和第 2 个特征(petal length)来画第一张图,是因为这两个特征的组合对三类样本的区分度最高。你可以自己换一下 x_idx 和 y_idx 的取值,比如改成 0 和 1,会看到 setosa 还是能分开,但 versicolor 和 virginica 几乎重叠在一起。这正好说明:不同特征组合的可视化效果差别很大,你在做数据分析时,不能只盯着某一张图下结论。
2.4 这个方案适合谁
我的建议是:零基础新手。手写一遍这个代码,你对 matplotlib 的画图流程、布尔索引、循环遍历类别这些操作都会有一个扎实的理解。等这些基础打牢了,再看后面几个方案,你会觉得它们只是“帮你封装好了细节”,而不是“黑魔法”。
3. 方案二:matplotlib 子图矩阵,一张图看完所有特征组合
3.1 方案一有一个明显的短板
方案一每次只能画两个特征的散点图,而鸢尾花数据集有 4 个特征,两两组合的话有 6 种情况。一个个地写 plt.scatter 太重复了,不仅代码冗余,看的时候还得手动来回切换对比。这时候可以用子图矩阵,把 4 个特征两两组合全部画在一张画布上,对角线上展示单个特征的分布,非对角线展示特征两两之间的关系。
3.2 用 subplots 构建 4x4 矩阵
python复制fig, axes = plt.subplots(4, 4, figsize=(16, 16))
for i in range(4):
for j in range(4):
ax = axes[i][j]
if i == j:
# 对角线:画该特征的直方图
ax.hist(iris.data[:, i], bins=15, color='#4ECDC4', alpha=0.7)
else:
# 非对角线:画类别散点图
for k, label in enumerate(iris.target_names):
mask = iris.target == k
ax.scatter(
iris.data[mask, j],
iris.data[mask, i],
c=colors[k],
s=20,
alpha=0.8,
)
if i == 3:
ax.set_xlabel(iris.feature_names[j], fontsize=8)
if j == 0:
ax.set_ylabel(iris.feature_names[i], fontsize=8)
ax.tick_params(labelsize=7)
plt.tight_layout()
plt.show()
3.3 为什么对角线要画直方图
对角线其实是同一个特征和自身组合,画散点图没有意义。改成直方图或者核密度曲线,可以展示单个特征在不同类别上的取值分布。比如 petal length 这个特征,setosa 的分布和其他两类基本不重叠,而 sepal width 上三类都有明显重叠。这些信息用一句话描述很抽象,但在子图矩阵里扫一眼就能得出直观结论。
3.4 解读这张大图的关键思路
看这张 4x4 的矩阵图,重点不是欣赏颜色,而是找“哪两个特征组合后,三个类别的点最不重叠”。我自己的经验是:优先看花瓣相关的组合,特别是 petal length 和 petal width 这个组合,三类点几乎可以看成三团互不相交的簇。这两个特征也是后面做分类模型时最重要的特征。
另外说一个小技巧:这个矩阵图如果每个子图都加 legend,会非常拥挤。所以我在代码里没有加图例,而是靠颜色区分。如果你需要在汇报时展示,可以单独给其中一个子图加图例,或者用 fig.legend() 统一加一个。
3.5 这个方案的定位
方案二适合做“数据探索阶段的总览图”。第一次拿到一个多特征数据集,先用它看全局,了解哪些特征有区分度、哪些特征存在异常分布,然后再用具体方案深入分析某一个组合。
4. 方案三:pandas 内置 scatter_matrix,一行代码快速出图
4.1 为什么不自己写矩阵了
如果你觉得上面那段 4x4 的循环代码还是太长了,pandas 里其实内置了画散点图矩阵的函数 scatter_matrix。它做的事情和方案二类似,但因为封装得好,只需要一两行代码就能实现。我自己在快速摸一个数据集底细的时候,经常用它先出一张全局图,再决定后续深入分析的方向。
4.2 核心代码
python复制from pandas.plotting import scatter_matrix
scatter_matrix(
df[iris.feature_names],
c=df['label'],
figsize=(12, 12),
alpha=0.8,
diagonal='kde',
)
plt.show()
注意导入方式:在新版 pandas 中,要用 from pandas.plotting import scatter_matrix,或者使用 pd.plotting.scatter_matrix(...)。如果你看到别人的老代码写的是 pd.scatter_matrix(...),在 pandas 1.0 以上版本会直接报 AttributeError,所以遇到报错先检查导入方式。
4.3 关键参数说明
c 参数传的是颜色映射的数值序列。这里直接传 df['label'],pandas 会根据数值大小自动映射到不同颜色,不需要你手动指定每个类别的颜色。diagonal='kde' 表示对角线画核密度估计曲线,比直方图更平滑一些。如果你偏执于直方图,改成 diagonal='hist' 就行。
alpha=0.8 控制透明度,数据量大的时候建议调低到 0.5 左右,避免点重叠严重导致看不清分布。figsize 控制画布大小,4 个特征就是 12 左右,特征再多的话要相应调大,否则子图会被压缩得很难看。
4.4 和方案二的区别
方案二自己写的矩阵图,自由度更高,比如可以对角线画密度图、非对角线画对数散点图、或者给特定子图加标注,这些都需要自己写逻辑。而 scatter_matrix 是快速输出,适合“先看一眼”,但不适合做精细定制的图。如果你只是想在分析过程中快速确认特征之间的关系,它完全够用;如果你要做最终汇报展示,建议用 seaborn 的方案。
4.5 在我实际使用中的体感
pandas 的 scatter_matrix 图风格比较朴素,颜色的默认组合也比较“工程风”。但它胜在不需要额外安装 seaborn,只要有了 pandas 和 matplotlib 就能画。在一个环境受限、不能随便 pip install 的情况下,这个方案能解决很多问题。
5. 方案四:seaborn 的 pairplot,颜值和功能都在线
5.1 为什么 seaborn 出图更好看
seaborn 是建立在 matplotlib 之上的高级可视化库,底层还是 matplotlib,但封装了大量统计图表的绘制逻辑。它处理分类着色、图例、分布曲线这些细节时,默认审美比 matplotlib 高不少,代码量却更少。pairplot 可以说是 seaborn 家族里画多特征总览图最方便的函数,没有之一。
5.2 核心代码
python复制import seaborn as sns
sns.set_theme(style='whitegrid')
sns.pairplot(
df[iris.feature_names + ['species']],
hue='species',
palette='deep',
)
plt.show()
只需要 3 行核心代码,就能得到一张漂亮的 4x4 网格图。
5.3 hue 参数到底做了什么
hue='species' 是 seaborn 的灵魂参数:指定用 species 这一列来给数据分组,并自动为不同类别使用不同颜色,同时生成图例。这就是“可视化三个类别”的关键——不需要手动循环三个类别去调颜色,一行 hue 参数全部搞定。
palette 控制配色方案,我常用的是 'deep'、'Set1'、'husl'。如果你有品牌色或者论文配色的要求,可以自定义一个颜色列表传进去,比如 palette=['#FF6B6B', '#4ECDC4', '#45B7D1']。
5.4 这张图比 scatter_matrix 多了什么
两者网格布局基本一致,区别主要体现在两点。
第一,pairplot 的对角线默认画的是直方图,并且会按照 hue 分组分别着色,你能直观看到每个类别在单个特征上的分布形态。如果用 diag_kind='kde',会变成核密度曲线,重叠区域看得更细腻。
第二,pairplot 可以配合 markers 参数同时使用形状和颜色双重编码:
python复制sns.pairplot(
df[iris.feature_names + ['species']],
hue='species',
markers=['o', 's', 'D'],
palette='deep',
)
这样即使打印成黑白图,也能通过点形状区分类别,对论文投稿或者打印场景很有用。
5.5 适合的场景
这是我目前最常推荐的方案。不管是自己分析数据,还是给同事演示特征分布,pairplot 出图效率高、效果好、信息量全。它的缺点是当特征数量较多时(比如超过 10 个),子图会变得非常小,图的内容会糊在一起。到那时候就不是简单画总览图能解决的了,需要做特征筛选或者降维。
6. 方案五:plotly 交互式可视化,探索数据和演示利器
6.1 前面的方案都是静态图,plotly 是完全不同的玩法
上面四个方案画的都是静态图片,图表固定在一个角度,你只能从画图前选好的特征组合去观察数据。而 plotly 画出的图是交互式的:可以鼠标拖拽旋转、缩放、悬停查看每个数据点对应的具体数值。在做探索性分析或者向别人展示数据时,这种交互体验是静态图完全比不了的。
6.2 用 plotly 画三维散点图
python复制import plotly.express as px
df_vis = df.copy()
df_vis['species'] = iris.target_names[iris.target]
fig = px.scatter_3d(
df_vis,
x='sepal length (cm)',
y='sepal width (cm)',
z='petal length (cm)',
color='species',
symbol='species',
opacity=0.8,
size_max=10,
)
fig.show()
运行后会在浏览器中打开一个可交互的 3D 散点图,三个轴分别对应三个特征,颜色和形状同时区分三个类别。鼠标拖拽就能旋转视角,悬浮在某个点上会显示它全部的特征值。
6.3 高维数据的新思路:PCA 降维后再可视化
鸢尾花数据集有 4 个特征,上面的 3D 图也只能展示其中 3 个。那剩下一个特征的信息就不看了吗?不一定。你可以用 PCA 把 4 维数据压缩到 2 维,再用二维散点图展示。PCA 的原理是把原始特征变换到新的正交坐标轴上,同时让新轴上的数据方差尽可能大,这样前两个主成分往往就能包含原始数据的大部分信息。
python复制from sklearn.decomposition import PCA
pca = PCA(n_components=2)
X_pca = pca.fit_transform(iris.data)
plt.figure(figsize=(8, 6))
for i, label in enumerate(iris.target_names):
mask = iris.target == i
plt.scatter(
X_pca[mask, 0],
X_pca[mask, 1],
c=colors[i],
label=label,
alpha=0.8,
edgecolors='white',
linewidths=0.5,
s=50,
)
plt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]:.2%})')
plt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]:.2%})')
plt.title('PCA on Iris Dataset (2D)')
plt.legend()
plt.grid(alpha=0.3)
plt.show()
explained_variance_ratio_ 表示每个主成分解释的方差比例。对鸢尾花数据集来说,前两个主成分通常能解释约 95% 以上的方差,也就是说从 4 维压到 2 维,损失的信息很少。PCA 降维后的结果里,三类样本也能比较清晰地分开,这为后面做分类模型提供了一个很直观的预判:这个数据集在低维空间里就具备较好的可分性。
6.4 使用 plotly 的注意点
plotly 第一次使用前需要安装:
bash复制pip install plotly
在 Jupyter Notebook 里运行 fig.show() 时,如果遇到图不显示的情况,大概率是渲染器配置的问题。可以手动指定渲染方式:
python复制import plotly.io as pio
pio.renderers.default = 'notebook'
如果你用的是 VS Code 之类的编辑器,也可以设置 pio.renderers.default = 'browser',图就会在默认浏览器中打开。另外,plotly 图形在某些离线环境下可能因为无法加载 plotly.js 而显示空白,这种情况比较少见,但遇到了可以检查网络和浏览器控制台的报错。
6.5 这个方案的定位
如果你的电脑配置不错,又需要做数据探索或者演示,plotly 是一个很好的选择。但如果你只是想把图表放进论文或者报告里,静态图就够用了。交互式的图表信息量虽然大,排版和印刷时反而不方便,这点需要结合自己的需求来判断。
7. 五个方案对比与选型建议
| 方案 | 代码量 | 交互性 | 信息维度 | 适合场景 | 依赖库 |
|---|---|---|---|---|---|
| 方案一:matplotlib 手写散点图 | 约 15 行 | 无 | 2 个特征组合 | 新手理解绘图原理、快速查看一对特征 | matplotlib |
| 方案二:matplotlib 子图矩阵 | 约 20 行 | 无 | 4x4 特征组合 | 数据探索阶段观察全局特征关系 | matplotlib |
| 方案三:pandas scatter_matrix | 1-2 行 | 无 | 4x4 特征组合 | 快速出概览图、环境依赖受限 | pandas |
| 方案四:seaborn pairplot | 1-2 行 | 无 | 4x4 特征组合+分布 | 汇报展示、美观优先、论文插图 | seaborn |
| 方案五:plotly 3D/交互式 | 3-5 行 | 有 | 3 特征可旋转 | 探索性分析、演示、交互式查看 | plotly |
选型逻辑我给你总结成四句话:
- 如果你是第一次接触散点图可视化,老老实实用方案一,把它彻底理解,后面所有方案对你来说都是锦上添花。
- 如果你拿到一份陌生数据集,想快速摸清特征之间的关系,方案四优先,信息量和代码量的性价比最高。
- 如果你环境受限装不了 seaborn 和 plotly,方案三用 pandas 就能顶上。
- 如果你要给别人现场演示数据或者做交互式分析,方案五最合适。
8. 常见问题与排查技巧实录
8.1 中文标签乱码或者显示成方框
matplotlib 默认字体不支持中文。如果你把标题、轴标签设成中文,输出的图里会出现很多小方框。解决方法是设置中文字体:
python复制plt.rcParams['font.sans-serif'] = ['SimHei']
plt.rcParams['axes.unicode_minus'] = False
Windows 上通常用 SimHei 或 Microsoft YaHei,macOS 可以用 PingFang SC,Linux 可以用 WenQuanYi Zen Hei。如果设置了还是方框,先用以下命令看系统里有哪些可用字体:
python复制import matplotlib.font_manager as fm
print([f.name for f in fm.fontManager.ttflist])
然后选一个中文字体名填进去。另外我建议在代码开头统一设置全局字体,不要每张图都写一遍。
8.2 scatter_matrix 报 AttributeError
如果你遇到这个报错:
text复制AttributeError: module 'pandas' has no attribute 'scatter_matrix'
说明你的 pandas 版本比较新,旧的调用方式被移除了。把代码改成:
python复制from pandas.plotting import scatter_matrix
scatter_matrix(df, ...)
或者:
python复制pd.plotting.scatter_matrix(df, ...)
两种方式都可以。遇到其他类似报错的,先怀疑是不是新版库的接口变动了,优先去官方文档查最近版本的变更说明。
8.3 seaborn 版本过老导致报名错误
某些老版本 seaborn 对 palette 参数的支持不太完善,或者 plot 风格函数名称已经变化。比如现在用 sns.set_theme(),而旧代码常见的是 sns.set()。这两种写法在一般场景下都能运行,但官方新版本推荐 set_theme()。如果你图出不来或者样式不生效,先检查一下 seaborn 版本:
bash复制pip show seaborn
版本低于 0.11 的话建议升级:
bash复制pip install -U seaborn
8.4 plotly 图在 Jupyter 不显示
最常见的原因是渲染器没有设置好。可以这样处理:
python复制import plotly.io as pio
pio.renderers.default = 'notebook'
或者在使用 fig.show() 时指定:
python复制fig.show(renderer='browser')
另外 plotly 依赖浏览器渲染,如果你在一台完全没有图形界面的服务器上运行,默认的渲染器可能是空白的。这种情况下可以把图保存成 HTML 文件:
python复制fig.write_html('iris_3d.html')
然后下载到本地浏览器打开,效果和交互式一样。
8.5 保存图片时颜色和图例不对
如果你用 plt.savefig() 保存图片后,发现颜色和图例和 plt.show() 显示的不一致,通常是因为没有指定 bbox_inches 导致图例被截断,或者 dpi 太低导致图片模糊。我常用的保存方式:
python复制plt.savefig('iris_scatter.png', dpi=300, bbox_inches='tight')
dpi=300 适合打印需求,bbox_inches='tight' 会自动调整画布,把所有内容都包括进去。注意保存图片最好在 plt.show() 之前执行,有些环境下 show() 会清空当前画布。
8.6 一个关于特征选择的经验之谈
我在指导新人做鸢尾花可视化时,经常看到有人随机选两个特征就画图,画完发现类别混在一起,就开始怀疑数据有问题。其实数据没问题,纯粹是特征组合选得不好。如果你不确定哪个特征组合有区分度,最快的办法是先画一张方案四的 pairplot 全局图,眼睛扫一遍再决定深入分析哪两个特征。盲目选特征画图,常常做了无用功。
一点额外的个人体会
做数据分析可视化的目的不只是“画一张好看的图”,而是通过图回答“数据长什么样”“类别之间能不能分开”“用哪些特征去分开”这些问题。我最初学鸢尾花的时候,也犯过上来就套模板、跑完什么也没看懂的毛病。后来踩了几次坑才意识到,可视化真正的价值在于观察和验证,而不是纯粹输出一张图。这几个方案里,最建议你手敲一遍的是方案一,它能帮你理解散点图的底层逻辑。日常分析中画图,我自己用最多的还是方案四,一行代码解决大部分需求。多试试不同方案,找到适合你自己工作流的那一套,比背下任何代码模板都管用。
