PyTorch入门学习(十四):优化器

目录

一、优化器的重要性

[二、PyTorch 中的深度学习](#二、PyTorch 中的深度学习)

三、优化器的选择


一、优化器的重要性

深度学习模型通常包含大量的参数,因此训练过程涉及到优化这些参数以减小损失函数的值。这个过程类似于找到函数的最小值,但由于模型通常非常复杂,所以需要依赖数值优化算法,即优化器。优化器的任务是调整模型参数,以最小化损失函数,从而提高模型的性能。

二、PyTorch 中的深度学习

PyTorch 是一个流行的深度学习框架,它提供了广泛的工具和库,用于创建、训练和部署深度学习模型。下面我们将通过一个简单的示例来了解如何使用 PyTorch 构建一个图像分类模型并训练它。

python 复制代码
import torch
import torchvision.datasets
from torch import nn
from torch.nn import Conv2d, MaxPool2d, Flatten, Linear, Sequential
from torch.utils.data import DataLoader

# 加载 CIFAR-10 数据集
dataset = torchvision.datasets.CIFAR10(root="D:\\Python_Project\\pytorch\\dataset2", train=False, transform=torchvision.transforms.ToTensor(), download=True)
dataloader = DataLoader(dataset=dataset, batch_size=1)

# 定义一个简单的神经网络模型
class Tudui(nn.Module):
    def __init__(self):
        super(Tudui, self).__init__()
        self.model1 = Sequential(
            Conv2d(3, 32, 5, padding=2),
            MaxPool2d(2),
            Conv2d(32, 32, 5, padding=2),
            MaxPool2d(2),
            Conv2d(32, 64, 5, padding=2),
            MaxPool2d(2),
            Flatten(),
            Linear(1024, 64),
            Linear(64, 10),
        )

    def forward(self, x):
        x = self.model1(x)
        return x

tudui = Tudui()

# 使用交叉熵损失函数
loss_cross = nn.CrossEntropyLoss()

# 使用随机梯度下降(SGD)优化器
optim = torch.optim.SGD(tudui.parameters(), lr=0.01)

# 训练模型
for epoch in range(20):
    running_loss = 0.0
    for data in dataloader:
        imgs, labels = data
        outputs = tudui(imgs)
        results = loss_cross(outputs, labels)
        optim.zero_grad()
        results.backward()
        optim.step()
        running_loss = running_loss + results
    print(running_loss)

在上面的代码中,使用 PyTorch 创建了一个名为 Tudui 的神经网络模型,并使用 CIFAR-10 数据集进行训练。在训练过程中,使用了随机梯度下降(SGD)作为优化器来调整模型的参数,以降低交叉熵损失函数的值。

三、优化器的选择

在深度学习中,有多种不同类型的优化器可供选择,每种都有其独特的特点。常见的优化器包括随机梯度下降(SGD)、Adam、RMSprop 等。选择合适的优化器通常取决于具体问题的性质和模型的结构。在上述示例中,使用了SGD,但可以根据需要尝试不同的优化器,以找到最适合的问题的那一个。

参考资料:

视频教程:PyTorch深度学习快速入门教程(绝对通俗易懂!)【小土堆】

相关推荐
m0_6770343515 分钟前
机器学习-异常检测
人工智能·深度学习·机器学习
张子夜 iiii32 分钟前
实战项目-----在图片 hua.png 中,用红色画出花的外部轮廓,用绿色画出其简化轮廓(ε=周长×0.005),并在同一窗口显示
人工智能·pytorch·python·opencv·计算机视觉
凯尔萨厮35 分钟前
Java学习笔记三(封装)
java·笔记·学习
YoungUpUp1 小时前
【文件快速搜索神器Everything】实用工具强推——文件快速搜索神器Everything详细图文下载安装教程 办公学习必备软件
学习·everything·文件搜索·实用办公软件·everything 工具·文件快速搜索·搜索神器
1373i1 小时前
【Python】pytorch安装(使用conda)
pytorch·python·conda
RaLi和夕1 小时前
单片机学习笔记.C51存储器类型含义及用法
笔记·单片机·学习
Niuguangshuo1 小时前
深度学习基本模块:Conv2D 二维卷积层
人工智能·深度学习
武文斌772 小时前
ARM工作模式、汇编学习
汇编·嵌入式硬件·学习·arm
lingggggaaaa2 小时前
小迪安全v2023学习笔记(八十讲)—— 中间件安全&WPS分析&Weblogic&Jenkins&Jetty&CVE
笔记·学习·安全·web安全·网络安全·中间件·wps
Jayyih3 小时前
嵌入式系统学习Day36(简单的网页制作)
学习·算法