5.Introduction to PyTorch YouTube Series--visualable train

前言

借助一些工具进行可视化训练,学习地址

训练

安装必要的工具

在自己的conda环境中进行安装

python 复制代码
pip install torch torchvision matplotlib tensorboard
  • torchvision:PyTorch 的计算机视觉工具库,提供常用数据集(如 CIFAR-10、ImageNet)、经典模型结构(如 ResNet)和图像预处理功能。
  • matplotlib:Python 最流行的绘图库,用于在训练过程中绘制损失曲线、查看样本图像等。
  • tensorboard:用于可视化训练指标、模型结构等。

数据预处理

  • 下载数据:从互联网获取 Fashion-MNIST 数据集
  • 预处理:将图片转换成张量格式
python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim

import torchvision
from torchvision.transforms import v2

import matplotlib.pyplot as plt
import numpy as np

from torch.utils.tensorboard import SummaryWriter

# 将数据集中的样本图像添加到TensorBoard:
# 将图像数据转为张量数据
transform = v2.Compose([
    v2.ToImage(),
    v2.ToDtype(torch.float32, scale=True),
    v2.Normalize((0.5,),(0.5,))
])
# 下载并划分数据集和验证集
training_set = torchvision.datasets.FashionMNIST(
    './data',
    download=True,
    train=True,
    transform=transform
    )
validation_set = torchvision.datasets.FashionMNIST(
    './data',
    download=True,
    train=False,
    transform=transform
)
training_loader = torch.utils.data.DataLoader(
    training_set,
    batch_size=4,
    shuffle=True,
    num_workers=2
)
validation_loader = torch.utils.data.DataLoader(validation_set,
                                            batch_size=4,
                                            shuffle=False,
                                            num_workers=2)
# 分类标签
classes = ('T-shirt/top', 'Trouser', 'Pullover', 'Dress', 'Coat',
        'Sandal', 'Shirt', 'Sneaker', 'Bag', 'Ankle Boot')

图片可视化

用日志器的add_image来写入到tensorboard中

  • 批次加载:用 DataLoader 把数据打包成小批量(batch),供后续训练时循环使用
  • 预览数据:用 matplotlib 把图片画出来,确认数据没问题
  • 创建日志记录器SummaryWriter
  • 记录第一张图片(add_image),验证 TensorBoard 能正常工作
  • 启动 TensorBoard 服务,在浏览器里查看日志内容
python 复制代码
def matplotlib_imshow(img, one_channel=False):
		# 张量转换成matplotlib可以显示的图像
    if one_channel:
        img = img.mean(dim=0)
    img = img / 2 + 0.5     # unnormalize
    npimg = img.numpy()
    if one_channel:
        plt.imshow(npimg, cmap="Greys")
    else:
        plt.imshow(np.transpose(npimg, (1, 2, 0)))

def t1():
    # 一批提取4个图像
    dataiter = iter(training_loader)
    images, labels = next(dataiter)

    # 查看下载好的图像
    img_grid = torchvision.utils.make_grid(images)
    matplotlib_imshow(img_grid, one_channel=True)
    # plt.show() 
    writer = SummaryWriter('runs/fashion_mnist_experiment_1')

    # 把图片数据写到日志文件里
    writer.add_image('Four Fashion-MNIST Images', img_grid)
    writer.flush()

执行t1后,执行命令tensorboard --logdir=runs,然后在浏览器中输入http://localhost:6006/查看,其中TensorBoard通过读runs/fashion_mnist_experiment_1 文件中写下的日志并在文件中展示,最后可以看到这样的结果

可视化训练

可以用add_scalaradd_scalars把训练过程中的损失等标量绘制成图

python 复制代码
class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 4 * 4, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 16 * 4 * 4)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

def t2():
    net = Net()
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)
    print(len(validation_loader))

    for epoch in range(1):  
        running_loss = 0.0

        for i, data in enumerate(training_loader, 0):
            # 每次取一个批次4张图
            inputs, labels = data
            optimizer.zero_grad()
            outputs = net(inputs)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()

            running_loss += loss.item()
            if i % 1000 == 999:    
                print(f'Batch {i + 1}')
                # 每1000个批次评估一下模型在验证集上的表现
                running_vloss = 0.0

                # 模型切换到评估模式
                net.eval() # 关闭梯度计算,节省内存并加速
                with torch.no_grad():
                    for j, vdata in enumerate(validation_loader, 0):
                        vinputs, vlabels = vdata
                        voutputs = net(vinputs)
                        vloss = criterion(voutputs, vlabels)
                        running_vloss += vloss.item()
                net.train() # 切回训练模式

                # 计算平均损失并写入tensorboard
                avg_loss = running_loss / 1000
                avg_vloss = running_vloss / len(validation_loader)

                writer.add_scalars('Training vs. Validation Loss',
                                { 'Training' : avg_loss, 'Validation' : avg_vloss },
                                epoch * len(training_loader) + i)

                running_loss = 0.0
    print('Finished Training')

    writer.flush()

训练结束后可以在tensorboard中看到两条曲线,横轴是训练步数(批次),纵轴是损失值,训练损失曲线会持续下降,验证损失曲线则反映模型在未见数据上的表现

模型结构可视化

add_graph()来写入,可以可视化模型的层结构。本质是取一个批次的数据,送给模型,然后追踪一下数据流。

python 复制代码
def t3():
    net = Net()
    # 取一个批次给模型,追踪一下数据流
    dataiter = iter(training_loader)
    images, labels = next(dataiter)

    writer.add_graph(net, images)
    writer.flush()

这个时候打开tensorboard,会出现graphs标签,点进去双击net就可以看到数据流向了

高维数据降维可视化

add_embedding把高维的图像特征投影到 3D 空间中,能直观地看到不同类别在特征空间中的分布情况。

这个功能主要用于检查特征学习效果,如果同类别的点聚在一起、不同类别的点分开,说明模型学到的特征是有区分度的,因此这个可视化应该在训练结束后进行观测

python 复制代码
# 数据集中随机抽取样本
def select_n_random(data, labels, n=100):
    assert len(data) == len(labels)

    perm = torch.randperm(len(data))
    return data[perm][:n], labels[perm][:n]

def t4():
    images, labels = select_n_random(training_set.data, training_set.targets)

    class_labels = [classes[label] for label in labels]

    # 展平图像,add_embedding对输入有要求
    features = images.view(-1, 28 * 28)
    writer.add_embedding(features,
                        metadata=class_labels,
                        label_img=images.unsqueeze(1))
    writer.flush()
    writer.close()

tensorboard中选择 PROJECTOR选项来查看

相关推荐
北京迅为43 分钟前
【迅为开发板专属工具】把烧写入口放进浏览器|Topeet RK Flash
linux·人工智能·嵌入式·rk3568·烧写
架构师汤师爷43 分钟前
DeepSeek Harness 暴涨 14.6万 Star,保姆级教程来啦~
人工智能
不一样的少年_1 小时前
我让 AI Agent 先别改代码,它怎么还是动手了?
人工智能·agent·ai编程
才聚PMP1 小时前
深陷技术内卷难突围?AI+项目管理开辟增值新赛道!
大数据·人工智能
阿里云大数据AI技术1 小时前
基于阿里云Milvus 构建电商图文智能搜索平台
人工智能
樊小肆1 小时前
DeepSeeker-Code源码导读04-上下文压缩
人工智能·agent
鲁邦通物联网1 小时前
充电站柔性负荷架构演进:网络卡顿导致设备烧损,如何依托边缘计算网关重塑本地调功防线?
人工智能·边缘计算·边缘计算网关·物联网网关·5g数采·边缘计算盒子·工业级边缘计算网关
weixin_471383031 小时前
17 Self-RAG —— 幻觉检测 + 答案质量评估
python·agent
叠层归一研究院1 小时前
如何用程序搭建一个 AGI 种子系统(三):生长如何对接物理与数学宇宙
人工智能·python·算法·机器学习·transformer·agi