pytorch 适合初学者 0基础学习

1.Dataset类

代码作用:把一个图片文件夹包装成 PyTorch 数据集,让你能查询图片数量,并按编号取出图片

标签。

python 复制代码
# 导入 PyTorch 的 Dataset 类,用来定义自己的数据集
from torch.utils.data import Dataset

# 导入图片处理工具 Image,用来打开和显示图片
from PIL import Image

# 导入 os 模块,用来拼接路径、读取文件夹中的名称
import os


# 定义自己的数据集类,类名为 MyData
# 括号中的 Dataset 表示 MyData 继承了 Dataset
class MyData(Dataset):

    # 初始化方法:创建 MyData 对象时,Python 会自动执行这个方法
    # self:表示当前创建的数据集对象,由 Python 自动传入
    # root_dir:根文件夹路径,例如 ".../train"
    # label_dir:类别文件夹名称,例如 "ants_image"
    def __init__(self, root_dir, label_dir):

        # 把根文件夹路径保存到当前对象中
        self.root_dir = root_dir

        # 把类别文件夹名称保存到当前对象中
        self.label_dir = label_dir

        # 拼接根路径和类别文件夹名称,得到图片所在文件夹的路径
        # 例如:"F:/zuo/pytorch/data/练手数据集/train/ants_image"
        self.path = os.path.join(root_dir, label_dir)

        # 获取该文件夹内的名称列表
        # 例如:["0013035.jpg", "另一张图片.jpg", ...]
        # 注意:这里存的是文件名,还没有读取图片内容
        # os.listdir 不保证排序,也会列出子文件夹和非图片文件
        self.img_path = os.listdir(self.path)

    # 取数据的方法:执行 ants_dataset[idx] 时会自动调用
    # idx 是索引;例如 idx=0 表示取列表中的第一项
    def __getitem__(self, idx):

        # 根据索引,从文件名列表中取出一个文件名
        # 例如:"0013035.jpg"
        img_name = self.img_path[idx]

        # 把图片文件夹路径与文件名拼接,得到这张图片的完整路径
        img_item_path = os.path.join(self.path, img_name)

        # 根据完整路径打开图片,得到 PIL 图片对象
        img = Image.open(img_item_path)

        # 使用类别文件夹名称作为标签
        # 当前数据集的标签都是字符串 "ants_image"
        # 标签来自文件夹名称,不是程序识别图片后得到的
        label = self.label_dir

        # 返回两个结果:图片对象和对应的标签
        return img, label

    # 获取数据集大小的方法:执行 len(ants_dataset) 时会自动调用
    def __len__(self):

        # 返回文件名列表中的元素数量
        # 如果文件夹中全是图片,这个数量就是图片数量
        return len(self.img_path)


# 从这里开始顶格,表示下面的代码不属于 MyData 类

# 创建一个 MyData 数据集对象,并保存到 ants_dataset 变量中
# 创建时会自动执行上面的 __init__ 方法
ants_dataset = MyData(
    "F:/zuo/pytorch/data/练手数据集/train",  # 传给 root_dir
    "ants_image"                          # 传给 label_dir
)

# len(ants_dataset) 会调用 __len__ 方法,得到数据集大小
# print() 把这个数量显示到终端
print(len(ants_dataset))

# ants_dataset[0] 会调用 __getitem__ 方法,此时 idx=0
# 方法返回一对结果,再分别赋给 img 和 label,这叫"解包"
# img 接收第一张图片,label 接收它的标签
img, label = ants_dataset[0]

# 打印标签,本例输出:ants_image
print(label)

# 调用系统的图片查看程序,显示取出的图片
img.show()
python 复制代码
from torch.utils.data import Dataset
from PIL import Image
import os
  • Dataset:PyTorch 的数据集基类,用来定义自己的数据集。
  • Image:Pillow 库中的图片工具,用来打开和显示图片。
  • os:这里用来拼接文件路径、读取文件夹中的文件名。
python 复制代码
class MyData(Dataset):

class 表示定义一个"类"。

MyData 起的类名,括号里的 Dataset 表示它继承了 PyTorch 的 Dataset 类

方法 作用 什么时候调用
__init__ 保存路径、获取文件名列表 创建 MyData(...)
__getitem__ 取出一张图片及其标签 使用 ants_dataset[0]
__len__ 返回数据集大小 使用 len(ants_dataset)

1.1:类的创建

当运行:

python 复制代码
ants_dataset = MyData(
    "F:/zuo/pytorch/data/练手数据集/train",
    "ants_image"
)

Python 会创建一个 MyData 对象,并执行:def init(self, root_dir, label_dir):

这时候里面的参数变成:

python 复制代码
root_dir = "F:/zuo/pytorch/data/练手数据集/train"
label_dir = "ants_image"

路径拼接

python 复制代码
self.path = os.path.join(root_dir, label_dir)

得到的路径指向:F:/zuo/pytorch/data/练手数据集/train/ants_image

os.listdir() 会列出该文件夹内的名称。

python 复制代码
self.img_path = os.listdir(self.path)
复制代码
self.img_path:这里存的是文件名列表,还没有打开图片。
python 复制代码
["ant1.jpg", "ant2.jpg", "ant3.jpg"]

1.2:获取图片数量

python 复制代码
print(len(ants_dataset))

运行该代码之后,会调用:然后会计算列表的数量

python 复制代码
def __len__(self):
    return len(self.img_path)

1.3:获取一张图片

python 复制代码
img, label = ants_dataset[0]

会调用:

复制代码
def __getitem__(self, idx):

首先获取文件名字:

复制代码
img_name = self.img_path[idx]

假设第一个文件名是 ant1.jpg,那么:

复制代码
img_name = "ant1.jpg"

拼出这张图片的完整路径:F:/zuo/pytorch/data/练手数据集/train/ants_image/ant1.jpg

复制代码
img_item_path = os.path.join(self.path, img_name)

打开图片:

复制代码
img = Image.open(img_item_path)

设置标签:标签来自你传入的文件夹名称,不是程序看了图片后识别出来的。 这个数据集中的所有图片都会得到同样的标签。1.

复制代码
label = self.label_dir

1.4:显示图片

复制代码
img.show()

2.TensorBoard

2.1:add_scalar

代码:模拟 20 个逐渐下降的 loss 数值,把它们记录到日志中,之后用 TensorBoard 查看曲线。它没有真正训练模型,只是在练习记录数据。loss 通常表示模型预测与目标之间的差距,训练时我们通常希望它逐渐减小。

python 复制代码
from torch.utils.tensorboard import SummaryWriter

# 日志将保存在 runs/hello 文件夹
writer = SummaryWriter("runs/hello")

# 模拟一个不断下降的 loss
for step in range(20):
    fake_loss = 1 / (step + 1)

    writer.add_scalar(
        "Loss/train",  # 曲线名称
        fake_loss,     # 纵坐标
        step           # 横坐标
    )

writer.close()
print("日志记录完成")
复制代码
writer = SummaryWriter("runs/hello")

这行创建了一个 SummaryWriter 对象,并用变量 writer 保存它。

"runs/hello" 是日志保存目录,属于相对路径,以程序运行时的当前工作目录为起点。

例如,当前工作目录是:

复制代码
F:/zuo/pytorch

日志就会保存在:

复制代码
F:/zuo/pytorch/runs/hello

日志目录不存在时,工具会创建它,里面通常会出现名称以 events.out.tfevents 开头的文件。

循环生成 20 个数据点

复制代码
for step in range(20):

range(20) 依次提供从 019 的整数,一共 20 个。

所以循环中的 step 会依次变成:

复制代码
0、1、2、3、......、19

step 是你起的变量名,这里表示"第几步"。它不会自动代表训练轮数,具体含义由记录数据的人决定。每次循环执行:

复制代码
fake_loss = 1 / (step + 1)

fake_loss 表示"模拟的损失值"。随着 step 增大,分母越来越大,结果越来越小:

把每一步的数值记录下来

复制代码
writer.add_scalar(
    "Loss/train",
    fake_loss,
    step
)

scalar 的意思是"标量",初学时可以理解为一个数值 ,例如 0.5

add_scalar() 这里的三个参数分别是:

复制代码
writer.add_scalar("曲线名称", 纵坐标数值, 横坐标步数)

对应到你的代码:

参数 当前内容 作用
曲线名称 "Loss/train" 标识这条曲线
纵坐标 fake_loss 本次记录的损失值
横坐标 step 本次记录的步数

关闭记录工具

复制代码
writer.close()

这行在循环外面,表示 20 个点全部记录完后再关闭。

它会把尚未写出的日志数据写入文件,并释放相关资源。

2.2:add_image

从硬盘读取一张蚂蚁图片,把它从 PIL 图片转换成 NumPy 数组,查看图片的数据结构,然后把图片写入 TensorBoard 日志,以便在 TensorBoard 网页中查看。

python 复制代码
import numpy as np
from PIL import Image
from torch.utils.tensorboard import SummaryWriter


image_path = (
    "F:/zuo/pytorch/data/练手数据集/"
    "train/ants_image/0013035.jpg"
)

# 1. 从硬盘打开图片
pil_image = Image.open(image_path).convert("RGB")

# 2. PIL 图片转换成 NumPy 数组
image_array = np.array(pil_image)

print("PIL尺寸:", pil_image.size)
print("数组形状:", image_array.shape)
print("数据类型:", image_array.dtype)
print("最小像素:", image_array.min())
print("最大像素:", image_array.max())

# 3. 创建日志记录器
writer = SummaryWriter("runs/image_demo")

# NumPy 图片通常使用 HWC 排列
writer.add_image(
    "Images/ant_numpy",
    image_array,
    global_step=0,
    dataformats="HWC"
)

writer.close()
print("图片日志记录完成")
复制代码
import numpy as np

导入 NumPy,取一个简短的名字 np。NumPy 是 Python 中专门处理大量数字和数组的工具库。

NumPy 主要用来处理数组。图片本质上也可以用一组数字表示,经常使用 NumPy 处理图片。

复制代码
from PIL import Image

从 Pillow 库中导入 Image

它负责打开图片:

复制代码
Image.open(image_path)

PIL 可以完成:

  • 打开图片
  • 查看图片尺寸
  • 裁剪、缩放、旋转图片
  • 转换颜色格式
  • 保存图片
python 复制代码
from torch.utils.tensorboard import SummaryWriter

导入 PyTorch 提供的 TensorBoard 日志记录工具。

SummaryWriter 可以记录很多内容,例如:

  • 一个数值:add_scalar()

  • 一张图片:add_image()

  • 多张图片:add_images()

  • 模型结构:add_graph()

  • 参数分布:add_histogram()

    pil_image = Image.open(image_path).convert("RGB")

这一行连续做了两件事。

复制代码
Image.open(image_path)

它会根据 image_path 找到图片,并把它作为一个 PIL 图片对象打开。

复制代码
.convert("RGB")

把图片转换为标准 RGB 彩色图片。

复制代码
image_array = np.array(pil_image)

这行把 PIL 图片转换成 NumPy 数组。

假设图片高度为 512、宽度为 768,那么它的数组结构大致是:

复制代码
512 行 × 768 列 × 3 个颜色通道

print("PIL尺寸:", pil_image.size)

pil_image.size 返回 PIL 图片的尺寸:(宽度, 高度)

复制代码
print("数组形状:", image_array.shape)

shape 表示数组在每个方向上的长度:(高度, 宽度, 通道数)

表示方式 顺序
PIL 的 size (宽度, 高度)
NumPy 的 shape (高度, 宽度, 通道)
PyTorch 图片 Tensor (通道, 高度, 宽度)
复制代码
print("数据类型:", image_array.dtype)

dtype 是 data type 的缩写,表示数组元素的数据类型。

普通 RGB 图片一般会输出:数据类型uint8

复制代码
print("最小像素:", image_array.min())

min() 会寻找整个数组中最小的数字。

复制代码
writer.add_image(
    "Images/ant_numpy",
    image_array,
    global_step=0,
    dataformats="HWC"
)

把图片写入日志

参数一:这张图片在 TensorBoard 中显示的标签名称。

参数二: TensorBoard 中显示的名称。

参数三:前面转换得到的 NumPy 数组

参数四:这张图片属于训练的第几步,作为记号

参数五:image_array 中的三个维度按照"高度、宽度、通道"排列。

3.Transforms

transforms 可以理解成一条"图片加工流水线":

复制代码
磁盘中的原始图片
        ↓
调整尺寸、裁剪、翻转等
        ↓
转换成 PyTorch Tensor
        ↓
归一化
        ↓
送给神经网络

从硬盘读取的图片,通常不能直接交给神经网络。图片大小,类型等等不一样。

因此,我们经常需要完成这些处理:

  • 把图片统一为相同大小
  • 把 PIL 图片转换成 Tensor
  • 调整像素值的范围
  • 对训练图片做随机变化,增加数据多样性
  • 根据训练需求进行归一化

这些操作统称为 transform,也就是"变换"。

一般是以下操作过程

复制代码
# 1. 定义加工规则
transform = transforms.Compose([
    操作1,
    操作2,
    操作3
])

# 2. 读取原始图片
img = Image.open(...)

# 3. 执行加工
img = transform(img)

# 4. 把处理后的 Tensor 交给模型
output = model(img)

3.1:基础操作

python 复制代码
from PIL import Image
from torchvision import transforms

img = Image.open("data/练手数据集/train/ants_image/6240338_93729615ec.jpg").convert("RGB")

to_tensor = transforms.ToTensor()
tensor_img = to_tensor(img)

print(type(img))
print(type(tensor_img))
print(tensor_img.shape)
复制代码
to_tensor = transforms.ToTensor()

表示创建一个图片转换工具。

复制代码
tensor_img = to_tensor(img)

表示把图片交给这个工具,获得转换结果。

PIL 图片通常可以理解为:HWC

转换成 Tensor 后,排列变为:通道 × 高度 × 宽度:C × H × W

3.2:归一化

归一化的基本公式是:

复制代码
新像素值 =(原像素值 - mean)/ std

transforms.Normalize(
    mean=[0.5, 0.5, 0.5],
    std=[0.5, 0.5, 0.5]
)

RGB 有三个通道,所以 meanstd 分别有三个数字:

复制代码
R 通道:(R - 0.5) / 0.5
G 通道:(G - 0.5) / 0.5
B 通道:(B - 0.5) / 0.5

3.3:组合

实际项目通常需要连续执行多个操作,比如:

  1. 修改图片大小
  2. 转成 Tensor
  3. 归一化

如果每次都单独写,会比较麻烦:

复制代码
img = resize(img)
img = to_tensor(img)
img = normalize(img)

因此 torchvision 提供了 Compose,用来把多个 transform 组成一条流水线:

python 复制代码
from torchvision import transforms

transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(
        mean=[0.5, 0.5, 0.5],
        std=[0.5, 0.5, 0.5]
    )
])

使用时只需:

复制代码
img = Image.open("ant.jpg").convert("RGB")
img = transform(img)

3.4:完整的具体例子

python 复制代码
# 从 PIL 库中导入 Image,用于读取图片
from PIL import Image
# 从 torchvision 中导入 transforms,用于处理图片
from torchvision import transforms
# --------------------------------------------------
# 1. 设置图片路径
# --------------------------------------------------

# 小括号中的两个字符串会被 Python 自动连接起来
image_path = (
    "F:/zuo/pytorch/data/练手数据集/"
    "train/ants_image/0013035.jpg"
)
# -------------------------------------------------
# 2. 从硬盘读取图片
# --------------------------------------------------

# Image.open():根据路径打开图片
# convert("RGB"):保证图片是 RGB 三通道彩色图片
pil_image = Image.open(image_path).convert("RGB")

# 查看处理前的图片信息
print("处理前的类型:", type(pil_image))
# PIL 图片的 size 按照(宽度,高度)排列
print("处理前的尺寸:", pil_image.size)


# --------------------------------------------------
# 3. 定义图片处理流水线
# --------------------------------------------------

# Compose 可以把多个图片处理操作组合起来
# 执行时会按照列表中从上到下的顺序进行处理
transform = transforms.Compose([

    # 将图片调整为固定大小
    # 参数顺序是(高度,宽度)
    transforms.Resize((224, 224)),

    # 以 50% 的概率对图片进行水平翻转
    # p=0.5 表示翻转概率是 50%
    transforms.RandomHorizontalFlip(p=0.5),

    # 把 PIL 图片转换为 PyTorch Tensor
    # 图片形状会从 HWC 形式转换成 CHW 形式
    # 像素通常也会从 0~255 缩放到 0.0~1.0
    transforms.ToTensor(),

    # 对图片的 R、G、B 三个通道分别进行归一化
    # 计算公式:新值 =(原值 - mean)/ std
    # 这里会将像素范围大致从 [0, 1] 变成 [-1, 1]
    transforms.Normalize(
        mean=[0.5, 0.5, 0.5],  # R、G、B 三个通道的均值
        std=[0.5, 0.5, 0.5]    # R、G、B 三个通道的标准差
    )
])


# --------------------------------------------------
# 4. 对图片执行 transforms
# --------------------------------------------------

# 把原始 PIL 图片传入 transform
# Compose 中的操作会按照从上到下的顺序依次执行
tensor_image = transform(pil_image)


# --------------------------------------------------
# 5. 查看处理后的图片信息
# --------------------------------------------------

# 处理后的图片已经变成 PyTorch Tensor
print("处理后的类型:", type(tensor_image))

# Tensor 图片的形状按照 [通道数, 高度, 宽度] 排列
# RGB 图片通常输出 torch.Size([3, 224, 224])
print("处理后的形状:", tensor_image.shape)

# 查看 Tensor 中元素的数据类型
# 一般是 torch.float32
print("处理后的数据类型:", tensor_image.dtype)

# 查看归一化后的最小值
# .min() 得到一个只有一个元素的 Tensor
# .item() 把这个 Tensor 转换成普通 Python 数字
print("处理后的最小值:", tensor_image.min().item())

# 查看归一化后的最大值
print("处理后的最大值:", tensor_image.max().item())

4:torchvision

对CIFAR10测试数据进行操作

python 复制代码
# 导入 torchvision
# torchvision 是 PyTorch 中专门处理图片、视觉数据集和视觉模型的工具包
import torchvision
# 从 torchvision 中导入 transforms 模块
# transforms 用于对图片进行转换和预处理
from torchvision import transforms
# ============================================================
# 1. 定义图片预处理方法
# ============================================================
# Compose 的作用是把多个图片处理步骤组合起来
# 以后可以在列表中继续添加 Resize、Normalize、随机翻转等操作
dataset_transform = transforms.Compose([
    # 将图片转换成 PyTorch 的 Tensor
    #
    # CIFAR10 原始图片通常是 PIL 图片
    # 转换前:
    #     PIL.Image.Image
    #
    # 转换后:
    #     torch.Tensor
    #
    # CIFAR10 图片转换后的形状为:
    #     [3, 32, 32]
    #
    # 3:RGB 三个颜色通道
    # 32:图片高度
    # 32:图片宽度
    #
    # ToTensor() 通常还会将像素范围:
    #     0~255
    # 转换为:
    #     0.0~1.0
    transforms.ToTensor()
])


# ============================================================
# 2. 设置数据集保存位置
# ============================================================

# CIFAR10 下载后会保存在这个文件夹中
#
# Windows 路径推荐使用正斜杠 /
# 这样可以避免反斜杠 \ 产生转义问题
dataset_root = "F:/zuo/pytorch/dataset"
# ============================================================
# 3. 下载并创建 CIFAR10 训练集
# ===========================================================
print("开始准备训练集......")
# 创建一个 CIFAR10 训练集对象
train_set = torchvision.datasets.CIFAR10(

    # 数据集保存的位置
    # 下载的数据会放入 F:/zuo/pytorch/dataset
    root=dataset_root,

    # train=True 表示使用训练集
    # CIFAR10 训练集一共有 50000 张图片
    train=True,

    # 指定图片预处理方法
    # 每次从 train_set 中取出图片时,
    # 都会自动执行前面定义的 ToTensor()
    transform=dataset_transform,

    # download=True 表示:
    #
    # 如果本地没有 CIFAR10,就自动下载;
    # 如果本地已经有完整的数据集,就不会重复下载
    download=True
)


# ============================================================
# 4. 下载并创建 CIFAR10 测试集
# ============================================================

print("开始准备测试集......")


# 创建一个 CIFAR10 测试集对象
test_set = torchvision.datasets.CIFAR10(

    # 训练集和测试集保存在同一个目录中
    root=dataset_root,

    # train=False 表示使用测试集
    # CIFAR10 测试集一共有 10000 张图片
    train=False,

    # 取出测试图片时,同样执行 ToTensor()
    transform=dataset_transform,

    # 检查数据是否存在
    # 如果不存在就自动下载
    download=True
)
# ============================================================
# 5. 查看数据集的基本信息
# ============================================================
print("下载并加载成功!")
# len(train_set) 得到训练集中的数据数量
# 正常情况下输出 50000
print("训练集数量:", len(train_set))
# len(test_set) 得到测试集中的数据数量
# 正常情况下输出 10000
print("测试集数量:", len(test_set))
# test_set.classes 保存了 CIFAR10 的所有类别名称
# 一共有 10 个类别
print("类别:", test_set.classes)

5:Dataloader取图片

python 复制代码
import torchvision
from torchvision import transforms
from torch.utils.data import DataLoader
# ============================================================
# 1. 定义图片转换规则
# ============================================================
dataset_transform = transforms.Compose([
    # 把 PIL 图片转换成 Tensor
    transforms.ToTensor()
])
# ============================================================
# 2. 创建训练集和测试集
# ============================================================
train_set = torchvision.datasets.CIFAR10(
    root="F:/zuo/pytorch/dataset",
    train=True,
    transform=dataset_transform,
    download=True
)
test_set = torchvision.datasets.CIFAR10(
    root="F:/zuo/pytorch/dataset",
    train=False,
    transform=dataset_transform,
    download=True
)
# ============================================================
# 3. 创建训练集 DataLoader
# ============================================================

train_loader = DataLoader(
    dataset=train_set,   # 从训练集中取数据
    batch_size=64,       # 每批取64张图片
    shuffle=True,        # 每轮训练前打乱顺序
    num_workers=0,       # Windows初学阶段使用0
    drop_last=False      # 保留最后不足64张的批次
)
# ============================================================
# 4. 创建测试集 DataLoader
# ============================================================

test_loader = DataLoader(
    dataset=test_set,    # 从测试集中取数据
    batch_size=64,       # 每批取64张图片
    shuffle=False,       # 测试时通常不打乱
    num_workers=0,
    drop_last=False
)
# ============================================================
# 5. 查看数据集和DataLoader的长度
# ============================================================
# Dataset 的长度表示图片总数
print("训练集图片数量:", len(train_set))
print("测试集图片数量:", len(test_set))
# DataLoader 的长度表示批次数量
print("训练集批次数量:", len(train_loader))
print("测试集批次数量:", len(test_loader))
# ============================================================
# 6. 取出训练集的第一批数据
# ============================================================
images, labels = next(iter(train_loader))
print("一批图片的形状:", images.shape)
print("一批标签的形状:", labels.shape)
print("一批标签:", labels)
# ============================================================
# 7. 查看前5张图片的类别名称
# ============================================================
for i in range(5):
    # labels[i] 是一个只有一个数字的Tensor
    # .item() 将它转换成普通Python整数
    label_number = labels[i].item()
    # 根据数字标签查找类别名称
    class_name = train_set.classes[label_number]
    print(
        "批次中的下标:", i,
        "数字标签:", label_number,
        "类别名称:", class_name
    )

dataset一次只能取出来一张图片

dataloader可以一次取多张

复制代码
train_loader = DataLoader(...)

DataLoader 是一个类

DataLoader(...):创建对象

复制代码
DataLoader(
    dataset=数据集,
    batch_size=每批数量,
    shuffle=是否打乱
)

dataset=train_set:表示 DataLoader 从哪个数据集中取数据。

batch_size=64:表示每次取出64条数据:图片+标签

num_workers=0:表示由多少个子进程负责读取数据。

drop_last:假设有50000张图片,50000 ÷ 64 = 781批,还剩16张,如果是false,则保留。

从 DataLoader 中取出一批数据

python 复制代码
# 创建一个迭代器
data_iterator = iter(train_loader)

# 从迭代器中取出第一批数据
images, labels = next(data_iterator)
复制代码
图片形状:torch.Size([64, 3, 32, 32])
标签形状:torch.Size([64])

[64, 3, 32, 32]
  ↑  ↑   ↑   ↑
  │  │   │   └── 宽度
  │  │   └────── 高度
  │  └────────── RGB三个通道
  └───────────── 这一批有64张图片
python 复制代码
for i in range(5):
    # labels[i] 是一个只有一个数字的Tensor
    # .item() 将它转换成普通Python整数
    label_number = labels[i].item()

    # 根据数字标签查找类别名称
    class_name = train_set.classes[label_number]

    print(
        "批次中的下标:", i,
        "数字标签:", label_number,
        "类别名称:", class_name
    )

CIFAR10 数据集对象内部保存了一个类别列表:

复制代码
print(train_set.classes)

输出:

复制代码
[
    'airplane',
    'automobile',
    'bird',
    'cat',
    'deer',
    'dog',
    'frog',
    'horse',
    'ship',
    'truck'
]

这是一个普通的 Python 列表,每个位置对应一个数字标签:

复制代码
下标0 → airplane
下标1 → automobile
下标2 → bird
下标3 → cat
下标4 → deer
下标5 → dog
下标6 → frog
下标7 → horse
下标8 → ship
下标9 → truck

例如:

复制代码
class_name = train_set.classes[3]

结果是:

复制代码
cat

6:tensor张量

python 复制代码
x = torch.tensor(5)

这行代码的意思是:创建一个数值为 5 的 PyTorch 张量,并用变量 x 保存它。

可以把"张量(Tensor)"先理解为 PyTorch 用来保存数值和进行计算的数据对象。它既可以保存一个数,也可以保存一组数,甚至一张图片的数据。

PyTorch 模型主要使用 Tensor 进行计算。随着后续学习,你还会用它处理一批数据、进行 GPU 计算,以及在满足条件时计算梯度。

python 复制代码
# 一个单独的数:零维张量
a = torch.tensor(5)
print(a)        # tensor(5)
print(a.shape)  # torch.Size([])

# 一个列表,列表里有一个数:一维张量
b = torch.tensor([5])
print(b)        # tensor([5])
print(b.shape)  # torch.Size([1])

# 一个列表,列表里有三个数:一维张量
c = torch.tensor([5, 6, 7])
print(c)        # tensor([5, 6, 7])
print(c.shape)  # torch.Size([3])

torch.tensor() 会根据传入数据的结构创建对应形状的张量。

7:nn.module

是 PyTorch 中所有神经网络模型和网络层的基础类。用来规定"输入数据经过什么计算,再得到什么输出"的模型外壳。

python 复制代码
import torch
from torch import nn


# 定义一个"输入加1"的模型
class AddOne(nn.Module):

    # 创建模型对象时执行
    def __init__(self):
        # 初始化父类 nn.Module
        super().__init__()

    # 规定数据进入模型后怎样计算
    def forward(self, x):
        # 输入加1,然后返回
        return x + 1

# 创建模型对象
model = AddOne()
# 创建输入Tensor
x = torch.tensor(5)
# 把x交给模型
y = model(x)
print("输入:", x)
print("输出:", y)

nn 是 PyTorch 中与神经网络有关的模块。

复制代码
class AddOne(nn.Module):

这句话可以拆成:定义一个叫 AddOne 的类,它继承 PyTorch 的 nn.Module

python 复制代码
class       定义一个类
AddOne      类的名字
nn.Module   被继承的父类
复制代码
super().__init__()

先把父类 nn.Module 自带的模型管理功能初始化好。

forward() 规定:输入数据进入模型后,要按照什么顺序进行计算。

_**init_**forward 的分工

python 复制代码
class MyModel(nn.Module):

    def __init__(self):
        super().__init__()

        # 在这里定义需要使用的网络层
        self.layer = nn.Linear(1, 1)

    def forward(self, x):
        # 在这里规定数据怎样经过这些网络层
        x = self.layer(x)

        return x
复制代码
__init__:准备零件

forward:规定零件怎样工作
python 复制代码
class MyModel(nn.Module):

    def __init__(self):
        super().__init__()

        # 准备两个零件
        self.linear = nn.Linear(1, 1)
        self.relu = nn.ReLU()

    def forward(self, x):
        # 规定数据经过零件的顺序
        x = self.linear(x)
        x = self.relu(x)

        return x
相关推荐
Rocky Ding*1 小时前
【三年面试五年模拟】2026-09-06 拼多多 AI Agent研发岗秋招笔试4道算法题完整题解
论文阅读·人工智能·深度学习·机器学习·aigc·ai-native·拼多多
小柯南敲键盘2 小时前
跨马翻译:AI批量图片翻译工具,跨境电商视频字幕翻译与智能抠图一体搞定
人工智能·python·音视频
2601_962297252 小时前
Python里behave和pytest-bdd哪个更适合中大型项目?为什么?
python·bdd·行为驱动开发·behave·pytest-bdd
luckystar513~2 小时前
Geo + AI:【时空智能体】技术剖析
人工智能·ai·gis·geoai·空间智能体·时空智能体
吨吨ai2 小时前
2026年9月8日|GPT‑6 Astra + Codex:Pro 开发者的 AI Agent 工具链
人工智能·gpt
leoZ2312 小时前
2026-09-09-springboot-cloud-deploy-pitfalls
java·前端·javascript·vue.js·人工智能·spring boot·后端
dozenyaoyida3 小时前
AI与大模型新闻日报 | 2026-09-09
人工智能·ai·chatgpt·大模型·新闻
xqqxqxxq3 小时前
AI Agent学习:主动工具发现(李博杰《深入理解 AI Agent》4.8观后总结)
人工智能·学习
海上彼尚3 小时前
Cursor 模型的强度实测排行
前端·人工智能·后端