news 2026/9/11 20:01:15

pytorch 适合初学者 0基础学习

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
pytorch 适合初学者 0基础学习

1.Dataset类

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

标签。

# 导入 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()
from torch.utils.data import Dataset from PIL import Image import os
  • Dataset:PyTorch 的数据集基类,用来定义自己的数据集。
  • Image:Pillow 库中的图片工具,用来打开和显示图片。
  • os:这里用来拼接文件路径、读取文件夹中的文件名。
class MyData(Dataset):

class表示定义一个“类”。

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

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

1.1:类的创建

当运行:

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

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

这时候里面的参数变成:

root_dir = "F:/zuo/pytorch/data/练手数据集/train" label_dir = "ants_image"

路径拼接

self.path = os.path.join(root_dir, label_dir)

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

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

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

1.2:获取图片数量

print(len(ants_dataset))

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

def __len__(self): return len(self.img_path)

1.3:获取一张图片

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通常表示模型预测与目标之间的差距,训练时我们通常希望它逐渐减小。

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 网页中查看。

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 可以完成:

  • 打开图片
  • 查看图片尺寸
  • 裁剪、缩放、旋转图片
  • 转换颜色格式
  • 保存图片
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:基础操作

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 组成一条流水线:

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:完整的具体例子

# 从 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测试数据进行操作

# 导入 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取图片

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 中取出一批数据

# 创建一个迭代器 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张图片
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张量

x = torch.tensor(5)

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

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

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

# 一个单独的数:零维张量 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 中所有神经网络模型和网络层的基础类。用来规定“输入数据经过什么计算,再得到什么输出”的模型外壳。

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

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

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

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

__init__forward的分工

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:规定零件怎样工作
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
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 20:00:50

RAG私有知识库毕设实战:从文档切分到本地LLM问答全流程

简介:这是一套面向计算机专业本科生的高分毕业设计级RAG私有知识库智能问答系统实现方案,专为毕设实战、课程设计与深度学习项目练手打造,解决学生缺乏端到端AI应用开发经验的痛点。资源包含545个文件,主体为145个Python源码&…

作者头像 李华
网站建设 2026/9/11 20:00:36

SQLite3 学习笔记:数据库基础、SQL 语句与 C 语言 API 详解

数据库 1 . 数据库文件与普通文件区别: 1)普通文件对数据管理(增删改查)效率低 2)数据库对数据管理效率高,使用方便 2. 常用数据库: 1.关系型数据库: 将复杂的数据结构简化为二维表格形式 大型:Oracle、DB2 中型:MySql、SQLServer 小型:Sqlit…

作者头像 李华
网站建设 2026/9/11 19:58:35

招聘数据可视化:Python爬虫与MapReduce全链路实践

简介:基于Python爬虫与MapReduce分析的招聘信息大数据可视化系统,是一份高分毕业设计整套资料,面向软件工程、计算机科学、人工智能等专业的学生,解决招聘信息采集、分布式分析及可视化展示的综合问题。资源内含完整系统源码、部署…

作者头像 李华
网站建设 2026/9/11 19:58:33

PCA异常检测完全指南:重构误差原理、代码实现与应用

简介:基于Python与PCA的异常检测算法设计实现资源,面向机器学习初学者和数据分析从业者,帮助理解利用主成分分析识别数据中异常模式的方法。内容覆盖PCA标准化、协方差矩阵、特征值分解、主成分选择与数据重构等核心步骤,并结合异…

作者头像 李华