PyTorch可视化双雄:Matplotlib与Seaborn深度揭秘
一、引言

在深度学习的世界里,PyTorch 已成为众多开发者和研究人员的首选框架之一,它以其动态计算图和强大的 GPU 加速能力,助力我们构建出各种复杂且高效的模型。然而,模型的构建仅仅是开始,想要真正理解模型的行为、分析数据的特征,可视化起着至关重要的作用。
可视化就像是我们窥探模型内部世界的一扇窗户,通过它,我们能将抽象的数据和复杂的模型结构转化为直观的图像、图表,从而更清晰地洞察模型在训练过程中的表现,比如损失函数的下降趋势是否合理,准确率是否稳步提升;也能深入了解数据的分布特征,判断数据是否存在异常值或不均衡的情况 。这些信息对于我们优化模型、提高模型性能,乃至做出更准确的决策都有着不可估量的价值。
而 Matplotlib 和 Seaborn 这两个强大的可视化工具,在 PyTorch 的生态系统中占据着重要地位。Matplotlib 作为 Python 的核心绘图支持库,提供了丰富的绘图函数和方法,能够满足各种基本绘图需求,从简单的折线图、散点图,到复杂的多子图布局,都能轻松应对,就像是一位基本功扎实的 “万能画家”。Seaborn 则是基于 Matplotlib 构建的高级可视化库,它在美观度和统计可视化方面更胜一筹,预设了多种精美的主题和调色板,让绘制出的图表不仅准确传达信息,还极具视觉吸引力,仿佛是为图表穿上了一件件华丽的外衣,特别适合用于展示数据之间的统计关系和分布情况。接下来,就让我们深入探索这两个工具在 PyTorch 中的应用吧。
二、Matplotlib 技术详解
2.1 安装与导入
Matplotlib 的安装十分便捷,如果你使用的是 pip 包管理器,只需在命令行中输入 pip install matplotlib ,pip 便会自动从官方仓库下载并安装 Matplotlib 及其依赖项。倘若你使用的是 Anaconda 环境,那么在 Anaconda Prompt 中执行 conda install matplotlib ,conda 会帮你处理好所有依赖关系,确保安装过程顺利进行 。
安装完成后,在 Python 代码中导入 Matplotlib 的常用方式为 import matplotlib.pyplot as plt ,这里的 plt 是一个约定俗成的别名,方便我们后续调用 Matplotlib 的绘图函数。
2.2 核心组件剖析
Matplotlib 的绘图系统包含几个核心组件,理解它们对于灵活绘图至关重要。
Figure:可以看作是一张画布,是所有绘图元素的顶级容器,一个 Python 脚本中可以包含多个 Figure 对象,每个 Figure 都代表一个独立的窗口或页面。例如,当我们想要对比两组完全不同的数据可视化结果时,就可以分别创建两个 Figure 对象来展示。
Axes:它是 Figure 中的一个子区域,也就是我们实际绘制图形的区域,每个 Axes 都有自己独立的坐标系(x 轴和 y 轴) ,一个 Figure 可以包含多个 Axes,这些 Axes 可以是不同类型的图表,比如折线图、散点图等同时存在于一个 Figure 中。
Axis:是 Axes 的组成部分,负责处理坐标轴的刻度、标签等相关设置,每个 Axes 都有两个 Axis 对象,分别对应 x 轴和 y 轴。
通过下面的代码示例,我们可以更直观地看到它们之间的创建关系:
import matplotlib.pyplot as plt
# 创建一个Figure对象,figsize参数设置画布大小为宽8英寸,高6英寸
fig = plt.figure(figsize=(8, 6))
# 在Figure对象上添加一个Axes对象,111表示将画布划分为1行1列,当前Axes占据第1个位置
ax = fig.add_subplot(111)
# 绘制一条简单的折线,x为[1, 2, 3, 4],y为[1, 4, 9, 16]
ax.plot([1, 2, 3, 4], [1, 4, 9, 16])
# 显示图形
plt.show()
在这段代码中,我们首先创建了一个 Figure 对象,然后在这个 Figure 上添加了一个 Axes 对象,并在 Axes 上绘制了一条折线。通过这样的操作,我们就建立起了 Figure、Axes 和实际绘图之间的联系 。
2.3 基础绘图实操
绘制折线图
折线图常用于展示数据随时间或其他连续变量的变化趋势。在 Matplotlib 中,使用 plt.plot() 函数绘制折线图。例如,我们要展示一个简单的数学函数 \(y = x^2\) 在 \(x\) 取值为 1 到 5 时的变化情况:
import matplotlib.pyplot as plt
# x轴数据
x = [1, 2, 3, 4, 5]
# y轴数据,通过计算x的平方得到
y = [i ** 2 for i in x]
# 绘制折线图,设置线条颜色为蓝色(b),线型为实线(-),标记点为圆形(o)
plt.plot(x, y, 'bo-')
# 设置图表标题
plt.title('Simple Line Plot')
# 设置x轴标签
plt.xlabel('X Values')
# 设置y轴标签
plt.ylabel('Y = X^2')
# 显示图表
plt.show()
在 plt.plot() 函数中, 'bo-' 是一个格式化字符串, 'b' 表示蓝色, 'o' 表示圆形标记点, '-' 表示实线。通过这样的设置,我们可以清晰地看到折线图的样式。同时,通过 plt.title() 、 plt.xlabel() 和 plt.ylabel() 函数,我们为图表添加了标题和坐标轴标签,使图表更具可读性 。
绘制散点图
散点图主要用于展示两个变量之间的关系,观察数据点的分布情况。使用 plt.scatter() 函数来绘制散点图。假设我们有一组随机生成的数据,用来探索两个变量之间可能存在的关系:
import matplotlib.pyplot as plt
import numpy as np
# 生成100个服从标准正态分布的随机数作为x轴数据
x = np.random.randn(100)
# 生成100个服从标准正态分布的随机数作为y轴数据
y = np.random.randn(100)
# 绘制散点图,设置点的颜色为红色(r),透明度为0.5
plt.scatter(x, y, c='r', alpha=0.5)
# 设置图表标题
plt.title('Scatter Plot')
# 设置x轴标签
plt.xlabel('X')
# 设置y轴标签
plt.ylabel('Y')
# 显示图表
plt.show()
在这个例子中,我们使用 np.random.randn() 函数生成了两组随机数据,然后通过 plt.scatter() 函数将这些数据点绘制在图表上。 c='r' 设置了点的颜色为红色, alpha=0.5 则设置了点的透明度,使得图表在展示大量数据点时不会过于密集,便于观察数据分布 。
绘制柱状图
柱状图常用于比较不同类别之间的数据大小。通过 plt.bar() 函数绘制柱状图。比如,我们有不同水果的销量数据,想要直观地比较它们的销量:
import matplotlib.pyplot as plt
# 水果类别
fruits = ['Apple', 'Banana', 'Orange', 'Mango']
# 对应的销量
sales = [35, 28, 42, 15]
# 绘制柱状图,设置柱子的颜色为绿色(g)
plt.bar(fruits, sales, color='g')
# 设置图表标题
plt.title('Fruit Sales')
# 设置x轴标签
plt.xlabel('Fruits')
# 设置y轴标签
plt.ylabel('Sales Quantity')
# 显示图表
plt.show()
在这段代码中,我们将水果类别作为 x 轴的刻度标签,销量作为 y 轴的数据,通过 plt.bar() 函数绘制出柱状图。不同水果对应的柱子高度直观地展示了它们销量的差异, color='g' 将柱子颜色设置为绿色,使图表更加美观 。
2.4 高级绘图技巧
添加图例
当图表中存在多个数据系列时,图例可以帮助我们区分不同的数据。在绘图时,为每个数据系列设置一个 label 参数,然后调用 plt.legend() 函数添加图例。例如,我们在一个图表中同时绘制两条折线:
import matplotlib.pyplot as plt
# x轴数据
x = [1, 2, 3, 4, 5]
# 第一条折线的y轴数据
y1 = [1, 4, 9, 16, 25]
# 第二条折线的y轴数据
y2 = [1, 8, 27, 64, 125]
# 绘制第一条折线,设置标签为'Y = X^2'
plt.plot(x, y1, label='Y = X^2')
# 绘制第二条折线,设置标签为'Y = X^3'
plt.plot(x, y2, label='Y = X^3')
# 添加图例,loc='upper left'表示将图例放置在左上角
plt.legend(loc='upper left')
# 设置图表标题
plt.title('Multiple Lines Plot')
# 设置x轴标签
plt.xlabel('X Values')
# 设置y轴标签
plt.ylabel('Y Values')
# 显示图表
plt.show()
在这个例子中,我们为两条折线分别设置了 label ,然后通过 plt.legend() 函数添加了图例,并使用 loc='upper left' 参数将图例放置在左上角,这样我们就能清晰地区分两条折线所代表的数据系列 。
设置坐标轴范围和刻度
有时候,我们需要根据数据的特点手动设置坐标轴的范围和刻度,以使图表展示更加合理。使用 plt.xlim() 和 plt.ylim() 函数设置坐标轴范围, plt.xticks() 和 plt.yticks() 函数设置刻度。比如,我们有一组数据,想要限制 x 轴范围在 0 到 10,y 轴范围在 0 到 100,并自定义 x 轴刻度:
import matplotlib.pyplot as plt
# x轴数据
x = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
# y轴数据
y = [10, 20, 30, 40, 50, 60, 70, 80, 90, 100]
# 绘制折线图
plt.plot(x, y)
# 设置x轴范围为0到10
plt.xlim(0, 10)
# 设置y轴范围为0到100
plt.ylim(0, 100)
# 设置x轴刻度为1到10
plt.xticks(range(1, 11))
# 设置图表标题
plt.title('Axis Settings')
# 设置x轴标签
plt.xlabel('X')
# 设置y轴标签
plt.ylabel('Y')
# 显示图表
plt.show()
在这段代码中,通过 plt.xlim(0, 10) 和 plt.ylim(0, 100) 分别限制了 x 轴和 y 轴的范围, plt.xticks(range(1, 11)) 将 x 轴刻度设置为 1 到 10,这样图表就能更聚焦于我们关注的数据区间 。
添加文本注释
为了在图表中强调某些关键信息,我们可以添加文本注释。使用 plt.text() 函数添加普通文本注释, plt.annotate() 函数添加带箭头的注释。例如,我们在一个散点图中标记某个特殊的数据点:
import matplotlib.pyplot as plt
# x轴数据
x = [1, 2, 3, 4, 5]
# y轴数据
y = [5, 4, 6, 2, 7]
# 绘制散点图
plt.scatter(x, y)
# 添加文本注释,在x=3,y=6的位置添加注释'特殊点'
plt.text(3, 6, '特殊点')
# 添加带箭头的注释,注释内容为'重点关注',箭头指向x=4,y=2的点
plt.annotate('重点关注', xy=(4, 2), xytext=(4.5, 3), arrowprops=dict(facecolor='black', shrink=0.05))
# 设置图表标题
plt.title('Annotated Scatter Plot')
# 设置x轴标签
plt.xlabel('X')
# 设置y轴标签
plt.ylabel('Y')
# 显示图表
plt.show()
在这个例子中, plt.text(3, 6, '特殊点') 在坐标 (3, 6) 处添加了一个普通文本注释, plt.annotate('重点关注', xy=(4, 2), xytext=(4.5, 3), arrowprops=dict(facecolor='black', shrink=0.05)) 则在坐标 (4, 2) 处添加了一个带箭头的注释,箭头从注释文本指向数据点, arrowprops 参数设置了箭头的样式,使我们能够更清晰地突出图表中的关键信息 。通过这些高级绘图技巧,我们可以让 Matplotlib 绘制出的图表更加专业、信息传达更加准确。
三、Seaborn 技术详解
3.1 安装与导入
Seaborn 同样可以使用 pip 进行安装,在命令行中执行 pip install seaborn ,pip 会自动下载并安装 Seaborn 及其相关依赖 。若使用 Anaconda 环境,在 Anaconda Prompt 中输入 conda install seaborn 即可完成安装。
在 Python 代码中,通常使用 import seaborn as sns 来导入 Seaborn 库,这里的 sns 是常用的别名,方便后续调用 Seaborn 的各种绘图函数和方法 。例如:
import seaborn as sns
import matplotlib.pyplot as plt
通过这样的导入方式,我们就可以在代码中使用 Seaborn 强大的可视化功能,并且可以结合 Matplotlib 进行更细致的图表定制 。
3.2 与 Matplotlib 关系解读
Seaborn 是基于 Matplotlib 构建的高级数据可视化库,它就像是 Matplotlib 的 “高级定制版”。Matplotlib 提供了基础的绘图功能,是构建可视化的基石,而 Seaborn 则在 Matplotlib 的基础上进行了更高层次的封装,提供了更简洁、美观的 API,使得创建复杂且美观的统计图表变得更加容易 。
Seaborn 利用 Matplotlib 的底层绘图机制来创建图表,比如它同样依赖 Matplotlib 的 Figure 和 Axes 对象来构建图表的基本结构 。但 Seaborn 对这些对象的操作进行了简化和优化,例如在设置图表样式、颜色主题等方面,Seaborn 提供了更便捷的方法。以绘制简单的折线图为例,Matplotlib 可能需要较多的代码来设置线条样式、颜色、标记等,而 Seaborn 只需通过几个参数就能实现同样的效果,并且生成的图表在默认情况下就具有更好的视觉效果 。
在实际使用中,我们可以将两者结合起来。当 Seaborn 提供的功能无法满足我们对图表细节的定制需求时,就可以借助 Matplotlib 的 API 进行进一步的调整 。比如,在使用 Seaborn 绘制完一个箱线图后,如果我们想要修改坐标轴的刻度标签格式,就可以使用 Matplotlib 的 ax.set_xticklabels() 和 ax.set_yticklabels() 方法来实现 。
3.3 特色绘图类型
分类图表绘制
Seaborn 提供了丰富的函数来绘制分类图表,帮助我们分析不同类别数据之间的关系和差异。
- barplot:用于绘制柱状图,展示不同类别数据的统计值(如均值、总和等)。例如,我们有不同城市的人口数据,想要比较各城市人口数量:
import seaborn as sns
import matplotlib.pyplot as plt
import pandas as pd
# 构建数据
data = {'City': ['Beijing', 'Shanghai', 'Guangzhou', 'Shenzhen'],
'Population': [21893095, 24870895, 18676605, 17560061]}
df = pd.DataFrame(data)
# 绘制柱状图,设置柱子颜色为橙色
sns.barplot(x='City', y='Population', data=df, color='orange')
# 设置图表标题
plt.title('Population of Different Cities')
# 设置x轴标签
plt.xlabel('City')
# 设置y轴标签
plt.ylabel('Population')
# 显示图表
plt.show()
在这个例子中, sns.barplot() 函数根据城市类别绘制了对应的人口数量柱状图, color='orange' 设置了柱子的颜色,使得图表更加醒目 。
- boxplot:箱线图可以展示数据的分布情况,包括中位数、四分位数和异常值等信息。比如,我们有不同班级学生的考试成绩数据,想了解各班级成绩的分布差异:
import seaborn as sns
import matplotlib.pyplot as plt
import pandas as pd
# 构建数据
data = {'Class': ['Class1', 'Class1', 'Class1', 'Class2', 'Class2', 'Class2'],
'Score': [85, 78, 92, 88, 76, 95]}
df = pd.DataFrame(data)
# 绘制箱线图
sns.boxplot(x='Class', y='Score', data=df)
# 设置图表标题
plt.title('Exam Scores Distribution by Class')
# 设置x轴标签
plt.xlabel('Class')
# 设置y轴标签
plt.ylabel('Score')
# 显示图表
plt.show()
通过箱线图,我们可以直观地看到不同班级成绩的中位数、上下四分位数以及是否存在异常值,从而对各班级的成绩分布有一个全面的了解 。
分布图表绘制
分布图表能够帮助我们了解数据的分布特征。
- kdeplot:核密度估计图(KDE plot)用于估计数据的概率密度函数,展示数据的分布形状 。例如,我们有一组学生的身高数据,想了解身高的分布情况:
import seaborn as sns
import matplotlib.pyplot as plt
import numpy as np
# 生成模拟身高数据,服从正态分布,均值为170,标准差为10
height_data = np.random.normal(170, 10, 1000)
# 绘制核密度估计图,设置颜色为蓝色
sns.kdeplot(height_data, color='blue')
# 设置图表标题
plt.title('Height Distribution')
# 设置x轴标签
plt.xlabel('Height (cm)')
# 设置y轴标签
plt.ylabel('Density')
# 显示图表
plt.show()
在这个例子中, sns.kdeplot() 函数根据身高数据绘制出核密度估计曲线, color='blue' 设置了曲线颜色,通过曲线我们可以清晰地看到身高数据的分布情况,比如峰值位置(即身高出现频率最高的值)以及数据的分布范围 。
3.4 样式与调色板运用
样式设置
Seaborn 预设了多种主题样式,使我们可以轻松改变图表的外观风格 。通过 sns.set_style() 函数来设置样式,可选的样式有 'whitegrid' 、 'darkgrid' 、 'white' 、 'dark' 、 'ticks' 等 。例如,将样式设置为 'whitegrid' ,可以在白色背景上添加灰色网格线,使图表更具可读性:
import seaborn as sns
import matplotlib.pyplot as plt
import numpy as np
# 设置样式为'whitegrid'
sns.set_style('whitegrid')
# 生成模拟数据
x = np.linspace(0, 10, 100)
y = np.sin(x)
# 绘制折线图
plt.plot(x, y)
# 设置图表标题
plt.title('Sin Function Plot')
# 设置x轴标签
plt.xlabel('X')
# 设置y轴标签
plt.ylabel('Y = Sin(X)')
# 显示图表
plt.show()
在这段代码中, sns.set_style('whitegrid') 将图表样式设置为 'whitegrid' ,然后绘制的折线图就会呈现出这种样式,网格线可以帮助我们更准确地读取数据在坐标轴上的位置 。
调色板使用
调色板用于控制图表中颜色的选择和搭配 。Seaborn 提供了多种预定义的调色板,如 'deep' 、 'muted' 、 'pastel' 、 'bright' 、 'dark' 、 'colorblind' 等 。通过 sns.set_palette() 函数来应用调色板 。例如,我们使用 'pastel' 调色板绘制一个散点图:
import seaborn as sns
import matplotlib.pyplot as plt
import numpy as np
# 设置调色板为'pastel'
sns.set_palette('pastel')
# 生成模拟数据
x = np.random.randn(100)
y = np.random.randn(100)
# 绘制散点图
sns.scatterplot(x, y)
# 设置图表标题
plt.title('Scatter Plot with Pastel Palette')
# 设置x轴标签
plt.xlabel('X')
# 设置y轴标签
plt.ylabel('Y')
# 显示图表
plt.show()
在这个例子中, sns.set_palette('pastel') 将调色板设置为 'pastel' ,绘制的散点图中的点就会使用 'pastel' 调色板中的颜色,这些柔和的颜色使图表看起来更加美观、舒适,不同颜色的点可以更好地区分不同的数据系列或类别 。
四、应用案例展示
4.1 案例一:PyTorch 模型训练过程可视化
在深度学习模型的训练过程中,监控损失函数和准确率的变化是评估模型性能和训练进度的关键。接下来,我们通过一个简单的 PyTorch 图像分类模型训练示例,展示如何使用 Matplotlib 将损失函数和准确率的变化过程可视化。
首先,导入必要的库:
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
然后,定义数据预处理和数据加载器:
# 数据预处理
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
# 加载训练集
train_dataset = datasets.CIFAR10(root='./data', train=True,
download=True, transform=transform)
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64,
shuffle=True)
# 加载测试集
test_dataset = datasets.CIFAR10(root='./data', train=False,
download=True, transform=transform)
test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=64,
shuffle=False)
接着,定义一个简单的卷积神经网络模型:
class SimpleCNN(nn.Module):
def __init__(self):
super(SimpleCNN, self).__init__()
self.conv1 = nn.Conv2d(3, 16, kernel_size=3, padding=1)
self.relu1 = nn.ReLU()
self.pool1 = nn.MaxPool2d(2)
self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
self.relu2 = nn.ReLU()
self.pool2 = nn.MaxPool2d(2)
self.fc1 = nn.Linear(32 * 8 * 8, 128)
self.relu3 = nn.ReLU()
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
out = self.conv1(x)
out = self.relu1(out)
out = self.pool1(out)
out = self.conv2(out)
out = self.relu2(out)
out = self.pool2(out)
out = out.view(-1, 32 * 8 * 8)
out = self.fc1(out)
out = self.relu3(out)
out = self.fc2(out)
return out
model = SimpleCNN()
再定义损失函数和优化器:
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
接下来进行模型训练,并记录每一轮训练的损失值和准确率:
# 存储损失值和准确率
train_losses = []
train_accuracies = []
test_losses = []
test_accuracies = []
# 训练模型
num_epochs = 10
for epoch in range(num_epochs):
model.train()
running_loss = 0.0
correct = 0
total = 0
for i, (images, labels) in enumerate(train_loader):
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
_, predicted = torch.max(outputs.data, 1)
total += labels.size(0)
correct += (predicted == labels).sum().item()
train_loss = running_loss / len(train_loader)
train_accuracy = correct / total
train_losses.append(train_loss)
train_accuracies.append(train_accuracy)
# 在测试集上评估模型
model.eval()
running_loss = 0.0
correct = 0
total = 0
with torch.no_grad():
for images, labels in test_loader:
outputs = model(images)
loss = criterion(outputs, labels)
running_loss += loss.item()
_, predicted = torch.max(outputs.data, 1)
total += labels.size(0)
correct += (predicted == labels).sum().item()
test_loss = running_loss / len(test_loader)
test_accuracy = correct / total
test_losses.append(test_loss)
test_accuracies.append(test_accuracy)
print(f'Epoch {epoch + 1}/{num_epochs}, Train Loss: {train_loss:.4f}, Train Acc: {train_accuracy:.4f}, Test Loss: {test_loss:.4f}, Test Acc: {test_accuracy:.4f}')
最后,使用 Matplotlib 绘制损失函数和准确率的变化曲线:
# 绘制损失函数曲线
plt.figure(figsize=(12, 6))
plt.subplot(1, 2, 1)
plt.plot(range(1, num_epochs + 1), train_losses, label='Train Loss')
plt.plot(range(1, num_epochs + 1), test_losses, label='Test Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Loss Curve')
plt.legend()
# 绘制准确率曲线
plt.subplot(1, 2, 2)
plt.plot(range(1, num_epochs + 1), train_accuracies, label='Train Acc')
plt.plot(range(1, num_epochs + 1), test_accuracies, label='Test Acc')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.title('Accuracy Curve')
plt.legend()
plt.show()
通过这些可视化的曲线,我们可以清晰地看到随着训练轮数的增加,损失函数是如何下降的,以及模型在训练集和测试集上的准确率变化情况。如果训练损失持续下降,而测试损失却上升,可能意味着模型出现了过拟合;如果两者都没有明显下降趋势,可能需要调整模型结构、优化器参数或增加训练数据等 。
4.2 案例二:数据集探索性分析
在使用 PyTorch 进行深度学习任务之前,对数据集进行探索性分析是非常重要的一步,它可以帮助我们了解数据的特征、分布情况以及变量之间的关系。Seaborn 提供了丰富的函数和方法,能够方便地对数据进行可视化分析。
我们以经典的鸢尾花数据集为例,展示如何使用 Seaborn 对 PyTorch 处理的数据集进行分析。首先,导入必要的库并加载数据集:
import seaborn as sns
import matplotlib.pyplot as plt
import pandas as pd
from torchvision import datasets
# 加载鸢尾花数据集
iris = datasets.load_iris()
df = pd.DataFrame(data=iris.data, columns=iris.feature_names)
df['target'] = iris.target
绘制特征间关系图
使用 Seaborn 的 pairplot 函数可以绘制数据集中各特征之间的两两关系图,同时可以根据目标变量(这里是鸢尾花的类别)对数据点进行着色,以便更好地观察不同类别数据在各特征上的分布差异:
g = sns.pairplot(df, hue='target')
plt.show()
在生成的关系图矩阵中,对角线上的图是每个特征的直方图或核密度估计图,展示了该特征的分布情况;非对角线上的图是两个特征之间的散点图,通过不同颜色的点表示不同的鸢尾花类别 。从这些图中,我们可以直观地看出,花瓣长度和花瓣宽度这两个特征对于区分不同种类的鸢尾花可能具有较高的判别力,因为不同类别的数据点在这两个特征上的分布差异较为明显 。
绘制箱线图分析特征分布
通过箱线图可以查看每个特征在不同类别下的分布情况,包括中位数、四分位数和异常值等信息,帮助我们进一步了解数据的特征和分布特征:
plt.figure(figsize=(12, 6))
for i, feature in enumerate(df.columns[:-1]):
plt.subplot(2, 2, i + 1)
sns.boxplot(x='target', y=feature, data=df)
plt.title(f'{feature} Distribution by Class')
plt.tight_layout()
plt.show()
在上述代码中,我们为每个特征绘制了一个箱线图,x 轴表示鸢尾花的类别,y 轴表示特征值。通过这些箱线图,我们可以清晰地看到每个特征在不同类别下的分布范围、中位数位置以及是否存在异常值 。比如,从花瓣长度的箱线图中可以看出,Setosa 类鸢尾花的花瓣长度明显小于其他两类,而且分布范围相对较窄 。
通过这些基于 Seaborn 的探索性分析,我们能够更深入地了解数据集的特征和潜在规律,为后续的模型选择、超参数调整以及数据预处理策略制定提供有力的依据 。
五、总结与展望
Matplotlib 作为 Python 可视化领域的基石,以其高度的可定制性和广泛的适用性,赋予我们精确控制图表每一个细节的能力。从简单的基础图表绘制,到复杂的多子图布局设计,Matplotlib 凭借丰富的函数和方法,满足了各种绘图需求,特别适合那些对图表样式有精确要求,需要高度个性化定制的场景,比如科研项目中的数据可视化展示,能够根据不同期刊的要求,精细调整图表元素 。
Seaborn 则站在 Matplotlib 的肩膀上,专注于提升统计图表的绘制效率和美观程度。它预设的多种精美主题和强大的统计可视化功能,让我们能轻松创建出专业且富有表现力的统计图表,在探索性数据分析(EDA)中发挥着巨大作用,能够快速帮助我们洞察数据的分布和变量之间的关系 。
在 PyTorch 项目中,这两个工具都是我们不可或缺的得力助手。Matplotlib 可用于直观展示模型训练过程中的各种指标变化,帮助我们及时发现模型训练的问题,调整训练策略;Seaborn 则能在数据预处理阶段,通过各种统计图表,深入分析数据集的特征和分布,为后续的模型设计和训练提供有力依据 。
更多推荐
所有评论(0)