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) 依次提供从 0 到 19 的整数,一共 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 有三个通道,所以 mean 和 std 分别有三个数字:
R 通道:(R - 0.5) / 0.5
G 通道:(G - 0.5) / 0.5
B 通道:(B - 0.5) / 0.5
3.3:组合
实际项目通常需要连续执行多个操作,比如:
- 修改图片大小
- 转成 Tensor
- 归一化
如果每次都单独写,会比较麻烦:
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