al+大数据每日学习笔记32

2026年9月4日

今天了解和学习深度学习中PyTorch框架

在深度学习落地开发中,框架的选择、环境的搭建、基础网络组件的使用是入门核心。本文将聚焦PyTorch框架,从主流深度学习框架对比、PyTorch环境安装、数据集加载、网络搭建、优化器、激活函数六大核心模块展开讲解,同时附上可直接运行的MNIST手写数字识别实战代码,帮助新手快速掌握PyTorch基础开发流程。

一、深度学习主流框架横向对比

目前深度学习行业主流框架各有优劣,结合就业市场需求,PyTorch是当下科研、项目落地的首选框架,各框架核心特点如下:

  • Caffe:传统经典框架,无需编码,仅通过配置文件即可搭建神经网络;缺点极其明显,安装流程繁琐,不支持新型网络模型,近年已停止版本更新,基本被行业淘汰。

  • TensorFlow(谷歌) :工业级主流框架,1.x版本代码冗余、学习门槛高;2.x版本整合Keras简化编码,但存在版本不兼容问题,新旧项目迁移成本高。

  • Keras:基于TensorFlow的高层封装框架,核心优势是代码极简、上手快,但灵活性较低,不适合复杂模型定制开发。

  • PyTorch(Facebook) :当下最热门的深度学习框架,核心亮点是上手极简、代码灵活、生态完善,模板化开发门槛低,适配科研实验、项目落地、竞赛等绝大多数场景,也是企业招聘核心技术栈。

二、PyTorch环境安装核心要点

PyTorch安装分为CPU版本 和GPU版本,二者核心差异源于硬件算力结构,深度学习训练优先推荐GPU版本。

2.1 CPU与GPU硬件核心区别

  • CPU(中央处理器):通用运算单元,结构以控制单元、缓存为主(占比75%),运算单元ALU仅占25%,擅长逻辑处理,不适合大规模矩阵运算。

  • GPU(图像处理器):专用并行运算单元,90%结构为ALU运算单元,控制、缓存占比极低,擅长深度学习海量数据并行计算,是模型训练的核心硬件。

GPU核心参数:显存容量、显存频率、显存位宽,直接决定模型训练速度与可承载的数据量,PyTorch GPU版本仅支持英伟达显卡。

2.2 CUDA核心概念

CUDA是NVIDIA推出的GPU并行计算平台与编程模型,无需依赖图形API,可直接调用GPU硬件算力完成深度学习矩阵运算,是PyTorch GPU版本运行的必备依赖。

2.3 GPU版本PyTorch安装流程

  1. 命令行输入 nvidia-smi,查看显卡支持的最高CUDA版本;

  2. 安装英伟达显卡驱动,保证驱动适配硬件;

  3. 安装CUDA工具包,安装版本必须低于硬件支持的最高版本;

  4. 输入 nvcc -V 验证CUDA是否安装成功;

  5. 通过官方pip命令安装对应版本PyTorch。

三、PyTorch核心基础模块详解

3.1 数据集加载与批量处理

深度学习训练的核心是数据迭代,PyTorch提供内置数据集工具与数据加载器,高效完成数据预处理:

  • datasets.MNIST:内置手写数字数据集,自动下载、划分训练集、测试集,无需手动处理数据;

  • DataLoader :核心批处理工具,通过 batch_size 参数将数据集打包为固定批次,支持迭代训练,是模型训练的必备组件。

3.2 基础网络层

本文实战用到两大核心网络层,适配图像分类基础模型:

  • nn.Flatten:展平层,将28×28的二维图像像素矩阵,转化为一维线性数据,适配全连接层输入;

  • nn.Linear:全连接层,实现特征映射,可自定义输入、输出维度,多层堆叠构建深度网络。

3.3 梯度下降优化器

优化器的核心作用是通过梯度反向传播,更新模型参数、降低损失,主流算法分为3类基础梯度下降法,同时衍生多种优化算法:

  • BGD批量梯度下降:使用全量数据计算梯度,收敛稳定、精度高,但内存占用大、训练速度慢;

  • SGD随机梯度下降:单次仅用单个样本更新参数,速度快,但更新方向不稳定,易陷入局部最优;

  • Mini-batch小批量梯度下降:工业界主流方案,折中BGD与SGD,将数据分为小批次迭代,兼顾速度与精度;

  • 进阶优化器:Momentum、AdaGrad、RMSprop、Adam、AdamW(自适应学习率,适配绝大多数场景)。

3.4 激活函数与梯度问题

激活函数赋予网络非线性拟合能力,同时直接影响梯度传播效果,是深度网络训练的关键:

3.4.1 梯度消失与梯度爆炸

核心成因:反向传播链式求导的连乘效应。

  • 梯度消失:Sigmoid函数导数区间为0-0.25,多层连乘后梯度趋近于0,浅层网络参数无法更新,网络失效;

  • 梯度爆炸:多层激活函数导数大于1,连乘后梯度无限增大,参数震荡不收敛。

3.4.2 最优解决方案

使用ReLU激活函数 替代Sigmoid,公式:f(x)=max(0,x),正值区间导数为1,彻底解决梯度消失问题,是深度网络默认激活函数。

四、MNIST手写数字识别

以下代码核心知识点:数据集加载、DataLoader批处理、自定义全连接网络、ReLU激活函数、Adam优化器,可直接复制运行,适配CPU/GPU环境。

复制代码
# 导入核心库
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import matplotlib.pyplot as plt

# 1. 数据预处理与数据集加载
# 定义数据归一化变换
transform = transforms.Compose([
    transforms.ToTensor(),  # 转为张量
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST数据集标准化参数
])

# 加载训练集、测试集
train_data = datasets.MNIST(
    root="./data", train=True, download=True, transform=transform
)
test_data = datasets.MNIST(
    root="./data", train=False, download=True, transform=transform
)

# 2. 数据批处理(核心:DataLoader)
batch_size = 64
train_loader = DataLoader(train_data, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_data, batch_size=batch_size, shuffle=False)

# 3. 搭建神经网络模型(Flatten+多层Linear+ReLU激活函数)
class MNISTNet(nn.Module):
    def __init__(self):
        super(MNISTNet, self).__init__()
        # 网络结构:28*28输入 -> 128 -> 256 -> 10分类输出
        self.flatten = nn.Flatten()  # 展平层
        self.fc1 = nn.Linear(28*28, 128)
        self.fc2 = nn.Linear(128, 256)
        self.fc3 = nn.Linear(256, 10)
        self.relu = nn.ReLU()  # 解决梯度消失的激活函数

    def forward(self, x):
        x = self.flatten(x)
        x = self.relu(self.fc1(x))
        x = self.relu(self.fc2(x))
        x = self.fc3(x)
        return x

# 4. 初始化模型、损失函数、优化器
model = MNISTNet()
criterion = nn.CrossEntropyLoss()  # 交叉熵损失(分类任务专用)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)  # Adam自适应优化器

# 5. 模型训练
def train(epochs=5):
    model.train()
    for epoch in range(epochs):
        total_loss = 0
        for batch_img, batch_label in train_loader:
            # 前向传播
            output = model(batch_img)
            loss = criterion(output, batch_label)
            # 反向传播与参数更新
            optimizer.zero_grad()  # 清空梯度
            loss.backward()  # 反向传播
            optimizer.step()  # 更新参数
            total_loss += loss.item()
        # 打印每轮训练损失
        print(f"第{epoch+1}轮训练,平均损失值:{total_loss/len(train_loader):.4f}")

# 执行训练
if __name__ == "__main__":
    train()

五、总结

  1. 框架选型上,明确PyTorch轻量化、易上手的核心优势,是入门与实战首选框架;

  2. 环境搭建层面,掌握CPU/GPU算力差异、CUDA核心作用及GPU版本安装规范,为模型高效训练奠定基础;

  3. 核心原理层面,吃透数据加载、网络层搭建、优化器迭代、激活函数与梯度问题四大核心知识点,解决深度学习训练的基础原理问题;

  4. 实操层面,通过MNIST手写数字识别案例,串联所有知识点,实现从数据处理、模型搭建到训练迭代的完整闭环,为后续复杂模型学习铺垫基础。

相关推荐
小羊没烦恼!3 天前
微服务化的基石——持续集成
java·大数据·word·powerpoint·.net
一隅论数智3 天前
给AI一张“业务概念地图“:本体如何从哲学走向企业智能
大数据·人工智能·经验分享·笔记·学习·学习方法·政务
XiHongShi20163 天前
STM32F407 RTC定时器例程,建议保存
stm32·单片机·学习
爱吃苹果的日记本3 天前
离散数学第六课
学习·离散数学
尧炎科技3 天前
防潮抗变形,就选纯品梅花全桉多层板
大数据
程序员大阳3 天前
副队长大数据教程(5)--集群情况下虚拟机网络配置
大数据·集群·nat·网路
自由能燃气设备3 天前
商用全预混低氮冷凝锅炉免费方案vs付费方案对比+选型避坑指南
大数据·数据库·人工智能
科创致远3 天前
科创致远 ESOP 系统核心效能与实战价值展示
大数据·数据库·人工智能·精益工程
嘉立创FPC苗工3 天前
FPC与机器人的双向赋能,解锁智能装备进化新势能
大数据·人工智能·制造·fpc·电路板