news 2026/8/31 3:34:41

PyTorch与TensorFlow双框架实战:环境配置到MNIST识别

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch与TensorFlow双框架实战:环境配置到MNIST识别

之前在业务里同时维护两个算法项目时,我经常被“PyTorch 还是 TensorFlow”这个问题卡住。网上观点各执一词,有人说 PyTorch 学术生态无敌,有人说 TensorFlow 部署链路成熟。实际上,一个算法工程师如果只会其中一套,遇到工业落地或者复现最新论文时会非常吃亏。

本文不打算替你站队,而是把两套框架放到同一套学习路径里完整串一遍:从环境准备、核心 API 对比,到 MNIST 手写数字识别的完整实战,再给出安装失败、模型加载报错等高频问题的排查思路。无论你是刚入门深度学习,还是已经能跑通简单模型但想补齐另一套框架,这篇文章都能帮你省下不少查资料的时间。

1. 背景与核心概念

1.1 深度学习框架是什么,为什么要同时掌握两个

深度学习框架的本质,是对“张量运算 + 自动求导 + 神经网络模块”的统一封装。框架存在的意义,是让你把精力花在模型结构、数据处理、训练策略上,而不是每次训练都从零写一遍反向传播。

很多初学者容易陷入“二选一”的纠结,但真实工程场景往往不是单选题。你要复现一篇最新论文,开源实现大概率在 PyTorch 生态里;你要给已有业务线做模型推理服务,线上系统可能早就跑在 TensorFlow Serving 上。这种情况下,能够同时看懂、同时操作两套框架,会比单纯站队某一方更游刃有余。

1.2 PyTorch 与 TensorFlow 的设计哲学差异

PyTorch 由 Meta(原 Facebook)主导开发,核心设计是动态计算图,也就是“边运行边建图”。你写的前向代码会真实执行,PyTorch 在背后通过自动微分记录梯度。这种模式的好处是调试非常直观:你可以在任意一行打印中间张量,可以使用原生 Python 的ifforprint,甚至可以打一个断点进入pdb调试。

TensorFlow 由 Google 开发,经历了 1.x 静态图到 2.x 动态图的大转变。TensorFlow 2 默认也开启了 Eager Execution,但保留了tf.function这套“编译为静态图”的能力,适合对性能有高要求的训练和推理场景。同时,TensorFlow 的高层 API 也就是 Keras,SequentialFunctional接口能让开发者快速搭出常见网络。

1.3 两者并不是非此即彼

从能力边界来说,PyTorch 同样可以做部署,TensorFlow 同样适合研究和教学。很多大型项目实际上是混合使用:模型在 PyTorch 里训练和验证,训练完成后导出为 ONNX 格式,再接入 TensorRT 或 TensorFlow 推理链路。理解两套框架的 API 设计,反而能让你更清楚地判断,某个功能到底应该在哪套生态里完成。

2. 环境准备与版本说明

2.1 安装前的整体规划

环境冲突是深度学习新手最常见的问题。PyTorch 和 TensorFlow 依赖的 CUDA 组件、底层库并不完全一致,如果直接装进同一个 Python 环境,很容易出现“装好这个、另一个就无法 import”的情况。

最稳妥的做法是使用 Anaconda 分别创建两个独立环境:

pytorch_env -> Python 3.10 + PyTorch 2.x + torchvision tf_env -> Python 3.10 + TensorFlow 2.x

隔离之后,两个环境互不干扰,即使一个环境被装坏,也不影响另一个。本文示例基于 Windows 10/11 或 Ubuntu 20.04/22.04,Python 3.10。你的机器 CUDA 版本可能不同,安装命令会有差异,但排查思路是通用的。

2.2 PyTorch 环境搭建

先创建环境:

conda create -n pytorch_env python=3.10 -y conda activate pytorch_env

接着检查 GPU 驱动信息:

nvidia-smi

上半部分输出里的CUDA Version: 12.x指的是驱动能够支持的最高 CUDA 版本,安装 PyTorch 时可以选择等于或低于它的 CUDA 运行时版本。例如驱动支持 CUDA 12.1,你可以安装 cu121,也可以选更保守的 cu118。

PyTorch 官网会根据你的系统环境生成安装命令。以 cu121 为例:

pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121

如果官方源下载较慢,可以换用国内镜像源,例如清华源、阿里源。CPU 版本安装更简单:

pip install torch torchvision torchaudio

这里有一个值得注意的坑:创建环境时 Python 版本不要太新。某些 PyTorch 版本尚未适配最新的 Python 小版本,会导致 pip 找不到对应 wheel 包。目前 Python 3.9 到 3.11 在大多数 PyTorch 版本中都比较稳妥。

2.3 TensorFlow 环境搭建

TensorFlow 安装比 PyTorch 特殊一点,尤其是 Windows 平台。先创建独立环境:

conda create -n tf_env python=3.10 -y conda activate tf_env

CPU 版本:

pip install tensorflow-cpu

如果是在 Linux 上想获得 GPU 支持,可以安装:

pip install tensorflow

需要注意,从 TensorFlow 2.11 开始,Windows 原生 pip 不再直接提供 GPU 支持。想在 Windows 上使用 GPU 运行较新版本的 TensorFlow,比如 2.16、2.18 等,推荐通过 WSL2(Windows Subsystem for Linux)或 Docker 安装。如果你的项目必须用 Windows 原生环境,那么 TensorFlow 2.10 是最后支持原生 GPU 的版本。很多人装完 TensorFlow 后一运行就报 DLL 或 GPU 相关错误,基本都是这个原因。

Jetson 这类 ARM 嵌入式平台要特别注意:不要直接从 PyPI 安装torch,大概率没有对应架构的 wheel。Jetson 上通常使用 NVIDIA 针对 JetPack 系统编译的预编译包,安装前先确认 JetPack 版本与 CUDA 版本,再到 NVIDIA 官方资源中找匹配的.whl

2.4 安装后的检查清单

安装完成后,分别进入两个环境检查。

PyTorch 检查:

import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU only")

TensorFlow 检查:

import tensorflow as tf print(tf.__version__) print(tf.config.list_physical_devices('GPU'))

如果torch.cuda.is_available()返回False,优先检查驱动版本、PyTorch 对应的 CUDA 运行时版本是否匹配,以及当前虚拟环境里是否真的安装了 GPU 版本。

3. 核心 API 与编程范式对比

3.1 张量操作

两个框架都以张量为基本数据结构,只是类型名不同。PyTorch 是torch.Tensor,TensorFlow 是tf.Tensor,底层都支持 GPU 运算和自动微分。

下面分别创建两个张量:

# PyTorch import torch a = torch.tensor([1.0, 2.0, 3.0]) b = a * 2 print(b)
# TensorFlow import tensorflow as tf a = tf.constant([1.0, 2.0, 3.0]) b = a * 2 print(b)

在 PyTorch 中,默认设备是 CPU,要用 GPU 必须手动.to('cuda');TensorFlow 检测到 GPU 后会自动尝试使用 GPU。两者的哲学差异在这里也能看出来:PyTorch 把设备管理交给开发者,更显式;TensorFlow 更自动,但也容易让人忽略资源分配细节。

3.2 模型定义

PyTorch 推荐继承torch.nn.Module,并在__init__中定义子层,在forward中定义前向传播:

import torch.nn as nn class MyNet(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(128, 10) def forward(self, x): return self.fc(x)

TensorFlow 中既可以写函数式模型,也可以继承tf.keras.Model

import tensorflow as tf from tensorflow.keras import layers, models # 方式一:Sequential model = models.Sequential([ layers.Dense(128, activation='relu'), layers.Dense(10, activation='softmax') ]) # 方式二:子类化 class MyNet(tf.keras.Model): def __init__(self): super().__init__() self.fc1 = layers.Dense(128, activation='relu') self.fc2 = layers.Dense(10, activation='softmax') def call(self, x): x = self.fc1(x) return self.fc2(x)

子类化时,TensorFlow 的call方法对应了 PyTorch 的forward。调用方式都是model(x),只是内部逻辑分别走forwardcall

3.3 训练循环

PyTorch 的训练循环通常手动编写,非常透明:

optimizer.zero_grad() outputs = model(inputs) loss = loss_fn(outputs, labels) loss.backward() optimizer.step()

TensorFlow 2 使用tf.GradientTape()显式记录求导过程:

with tf.GradientTape() as tape: predictions = model(inputs) loss = loss_fn(labels, predictions) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables))

两者本质相同:前向传播、计算损失、反向传播、更新权重。区别在于 PyTorch 把求梯度隐藏在backward()内部,TensorFlow 则通过tape.gradient()把“求梯度”这一步显式交给你。

3.4 数据集与数据管道

PyTorch 的数据处理通常基于torch.utils.data.DatasetDataLoader,你可以自定义 Dataset 类,用迭代器批量取数。TensorFlow 推荐使用tf.data.Dataset,它能把文件读取、预处理、乱序、批处理、预取组合成一条完整的数据管道。

在中小型项目里,两套体系差别不大。但在大数据量或生产链路中,tf.data提供了更丰富的并行预处理能力;PyTorch 则更强调和 Python 生态的无缝衔接,方便做复杂的在线数据增强。

3.5 PyTorch 2.6 权重加载变化

从 PyTorch 2.6 开始,torch.loadweights_only参数默认值变成了True。这个改动是为了提升安全性,避免恶意 pickle 文件在反序列化时执行任意代码。

带来的直接影响是:一些旧代码直接torch.load("model.pth")会报WeightsUnpickler相关异常。如果你加载的是完全可信的模型权重,可以显式设置weights_only=False

state_dict = torch.load("model.pth", map_location="cpu", weights_only=False)

如果加载的是标准state_dict,默认的weights_only=True通常能够直接工作。遇到这类报错时,先判断模型文件来源,再决定是否关闭限制。

4. 完整实战——手写数字识别两种实现

这里用 MNIST 手写数字识别做例子。原因是它对算力要求很低,两套框架都内置了数据下载接口,加上代码量不大,正好可以同题对比。

4.1 项目结构

dl-stack/ ├── mnist_pytorch.py ├── mnist_tensorflow.py └── README.md

两个脚本互相独立,分别在自己的 conda 环境中运行。

4.2 PyTorch 版 CNN 完整代码

文件路径:mnist_pytorch.py

import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms # 1. 数据预处理 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform) test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=256, shuffle=False) # 2. 定义 CNN 模型 class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1) self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1) self.pool = nn.MaxPool2d(2, 2) self.relu = nn.ReLU() self.fc1 = nn.Linear(64 * 7 * 7, 128) self.fc2 = nn.Linear(128, 10) def forward(self, x): x = self.pool(self.relu(self.conv1(x))) # 28 -> 14 x = self.pool(self.relu(self.conv2(x))) # 14 -> 7 x = x.view(-1, 64 * 7 * 7) x = self.relu(self.fc1(x)) return self.fc2(x) model = CNN() device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) # 3. 训练 for epoch in range(5): model.train() running_loss = 0.0 for images, labels in train_loader: images, labels = images.to(device), labels.to
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/31 3:34:00

海康威视4G无线监控摄像机:从选型到部署全指南

做安防项目或给老家装监控时,最常被问到一个问题:“我那里没有网线,也没有WiFi,能不能装监控?”以前遇到这种需求,只能拉网线、扯插座,或者用无线网桥转发,成本高、维护也麻烦。现在…

作者头像 李华
网站建设 2026/8/31 3:33:54

SpringBoot+Vue宠物领养系统毕业设计:从架构到部署全指南

简介:这是一套面向计算机专业本科生的高质量毕业设计实战资源,聚焦宠物领养业务场景,采用Spring Boot(后端)与Vue.js(前端)主流技术栈构建全栈Web系统,适用于毕业设计、课程设计及期…

作者头像 李华
网站建设 2026/8/31 3:32:34

HyperMesh 3D模块零基础入门:核心原理与网格生成实践

很多第一次接触 HM 的朋友,打开界面后的第一反应不是“我该学哪个功能”,而是“这个软件的按钮怎么这么多”。尤其是切到 3D 模块之后,菜单、面板、参数、拓扑颜色、网格类型挤在一起,完全不知道从哪里下手。如果你也是零基础&…

作者头像 李华
网站建设 2026/8/31 3:29:48

三极管放大电路的偏置供电:静态工作点与直流偏置详解

看到标题里的“偏执供电”,估计不少刚接触三极管的同学会愣一下:这个词是不是写错了?严格来说,业内规范叫法是“偏置供电”,英文是 Bias。不过从问题上也能看出来,真正让人困惑的不是错别字,而是…

作者头像 李华
网站建设 2026/8/31 3:28:23

驱动安装完全指南:从USB转串口到设备管理器排错

开发过程中最让人头疼的事情之一,就是设备明明插上了,电脑却毫无反应,或者设备管理器里冒出一个带黄色感叹号的“未知设备”。尤其在嵌入式开发、单片机调试、Arduino/STM32 折腾、USB 转串口通信这些场景里,驱动安装几乎是绕不开…

作者头像 李华
网站建设 2026/8/31 3:27:58

黄仲贤的大双摇Fender是什么琴?演唱会吉他考证全解析

如果你在 B站 或视频号刷到 Beyond 2005 Live 的片段,经常能看到弹幕里有人在问:黄仲贤手里那把大双摇 Fender 到底是什么琴?这个问题乍看很简单,真往下挖就会发现,它牵出了 Fender 原厂双摇历史、改装市场、演唱会镜头…

作者头像 李华