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手写数字识别案例,串联所有知识点,实现从数据处理、模型搭建到训练迭代的完整闭环,为后续复杂模型学习铺垫基础。

相关推荐
IT研究室1 小时前
最新大数据毕业设计选题推荐-基于大数据的台北市住宅价格数据可视化分析-大数据-Spark-Hadoop-Bigdata
大数据·信息可视化·课程设计
xian_wwq1 小时前
【学习笔记】深度认知系列-第10讲AI绘画与设计——从Stable Diffusion到Midjourney
笔记·学习·ai作画
starzy19901 小时前
Flink基础之Flink应用场景及特点优势:四大场景与核心优势
大数据·flink
凯尔萨厮1 小时前
Java学习笔记十五(GUI)
java·笔记·学习
招财小梗1 小时前
沈阳AI企业咨询可定制数字化方案吗?
大数据·人工智能·python
其实防守也摸鱼1 小时前
教育信息技术应用创新---基础软件信息赛(题库)
大数据·运维·人工智能·web安全·自动化
陈年老古董1 小时前
矿物分类实战:从传统机器学习到深度学习(含PyTorch实现)
笔记·深度学习·学习·机器学习·分类
咖啡忍者1 小时前
【SAP】100小时学会SAP-5生产计划PP
笔记
春风解人意1 小时前
从零开始学习嵌入式P35----数据库
数据库·嵌入式硬件·学习