PyTorch、TensorFlow、JAX核心API对比与安装指南
2026/8/31 13:34:54 网站建设 项目流程

深度学习框架的核心 API,决定了你写的每一行训练代码长什么样。最近被问得最多的,不是哪个框架的算法更强,而是 PyTorch 安装、TensorFlow 2.18 安装、Anaconda 配环境这类问题。很多人相信只要环境跑通,后面的模型代码就顺了;但实际经验是,环境只是入口,真正决定你后面顺不顺手的,是框架的 API 设计。网上经常看到“七大框架对比”这样的标题,但把七个框架摊开写,很容易变成资料堆。我更愿意把“7”理解为七组核心 API:张量、自动微分、模型构建、训练循环、数据加载、设备与分布式、导出部署。把 PyTorch、TensorFlow、JAX 这三套主流框架放进同一个坐标系里看,远比挨个介绍框架更有价值。

1. 先理解三套 API 背后的设计哲学

1.1 命令式、声明式与函数式,是分水岭

三个框架表面上都是“张量 + 自动微分”,但底层执行模型完全不同。PyTorch 默认是命令式动态图,你写一行 Python,它就立刻执行一行,像写普通程序一样。这对调试极其友好,print直接能看到中间结果,断点想加就加。TensorFlow 则走过一条弯路:早期版本用静态计算图,先把整张图定义好,再放到 Session 里执行。这个设计提升了部署性能,却让初学者很难理解。到 TensorFlow 2.0 引入 Eager Execution,默认也变成了立刻执行,但很多老 API 和进化痕迹还是留了下来。JAX 则完全是另一套思路:它不是给模型设计高层类,而是把“计算函数”和“函数转换”作为核心,通过jax.gradjax.jitjax.vmap这类函数变换来组织代码。

这个区别为什么会直接影响 API?因为框架要想做到自动求导,必须能追踪计算过程。PyTorch 在张量上用requires_grad标记和动态图回溯;TensorFlow 用GradientTape记录正向计算过程;JAX 则要求你的损失函数是纯函数,然后对它做函数变换。你可以用差不多的数学公式写出同一个模型,但训练循环和调试方式会很不一样。最直接的感受是:PyTorch 像是给每个张量装了一个记录器,TensorFlow 像在旁边放了一卷录音带,JAX 更像是一个可以把“函数”整体变形成“导函数”的高阶函数工厂。

1.2 出身决定了 API 的性格

PyTorch 脱胎于 Torch,底层用 C 加速,上层用 Python 做交互。它从一开始就选择了“让研究人员先跑起来”的路线,所以 API 非常贴 Python 习惯,Moduleoptimizerdataloader这些对象都很直觉。TensorFlow 出身于 Google 的分布式计算环境,早期关注点是“大规模部署”和“生产链路”,所以它最强的不是上手体验,而是从训练到上线的一整套工程能力。Keras 作为高层接口,把模型定义、编译、训练、导出封装得很短,这是 TensorFlow 最值得用的部分。JAX 同样是 Google 出品,但它不是 TensorFlow 的替代,而是面向数值计算和高性能研究的一套底层工具。它没有自己的一套独立模型库,通常要配合 Flax、Haiku、Equinox 等外部库使用。

正因为出身不同,三套 API 的“体感”差异会被放大。PyTorch 默认把一切都暴露给你,自由度高,但你需要自己处理很多细节;TensorFlow 给你提供了 Keras 这个安全气囊,但一旦需要深度自定义,就会碰到封装层过厚的问题;JAX 把底层机制暴露得非常彻底,却也把组织代码的责任还给了你。了解这一点,再看后面的核心 API 对比,就不会觉得某个框架“奇怪”,而是会理解它为什么长成这样。

2. 七个核心 API 维度横向对比

在逐项拆解前,可以先看一张总表。这张表不追求覆盖每个函数的细节,只标出三者在同一件事上的入口差异。

API 维度PyTorchTensorFlowJAX
张量创建torch.tensor/torch.zerostf.constant/tf.Variablejnp.array/jnp.zeros
自动微分backward()动态图tf.GradientTape记录jax.grad函数变换
模型构建nn.Module类 +forwardtf.keras.Model/Sequential参数 PyTree + 外部库
训练循环手动 for 循环model.fit()或自定义循环@jax.jit的训练 step 函数
数据加载DataLoader/Datasettf.data.Dataset通常复用 TF/PyTorch 数据管道
设备与分布式.to(device)/ DDPtf.distribute.Strategyjax.devices()/pmap
导出部署TorchScript /torch.exportSavedModel / TFLite / Servingjax2tf或直接服务函数

这张表最直观的信息是:TensorFlow 在模型构建和数据加载上给了你现成的高层入口,PyTorch 把控制权交给你,JAX 则倾向于让你用“函数 + 参数结构”自行组合。下面逐个展开,不讨论每个函数的所有参数,只抓住最影响开发方式的部分。

2.1 张量创建与基础运算

PyTorch 的张量 API 从 Python 用户角度比较自然。torch.tensor([1, 2, 3])torch.zeros(3, 4)torch.randn(2, 3)几乎不需要解释。它的一个显著特点是大量支持 in-place 操作,比如x.add_(1)会直接修改x。这个设计让内存使用更高效,但也带来副作用,尤其在需要对梯度追踪时,in-place 操作可能改变历史计算图,是初学容易踩坑的点。

TensorFlow 的tf.constant创建的是不可变张量,tf.Variable才是可训练参数。这种区分比 PyTorch 更严格,但也提醒你:在 TensorFlow 里,模型的可变状态与普通常量是分开管理的。tf.Tensor.numpy()可以把张量转成 NumPy 数组,调试时很方便。JAX 则更彻底:jnp.array是不可变对象,没有 in-place 操作。这意味着你不能写arr += 1然后期待原数组变化,而应该写成arr = arr + 1。对已经习惯 NumPy 和 PyTorch 的人来说,一开始会觉得别扭,但这是 JAX 为了能够做函数编译和自动并行而必须付出的代价。

从工程角度,理解张量不可变性很重要。PyTorch 的张量是“对象”,带状态;TensorFlow 的tf.Variable是显式的可变状态;JAX 的数组是“值”,天然安全。底层机制决定了调试体验和并发安全边界。

2.2 自动微分 API

自动微分是框架最核心的部分。PyTorch 的写法通常是:

import torch x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) y = (x ** 2).sum() y.backward() print(x.grad) # tensor([2., 4., 6.])

requires_grad表示这个张量需要梯度,backward()触发反向传播,梯度累积到x.grad。这套 API 直观,但也隐含着动态图的代价:每次前向计算都会构建一张图,如果你想在with torch.no_grad():之外做纯推理,就必须注意关闭梯度,否则会浪费内存。

TensorFlow 的自动微分用GradientTape显式记录正向过程:

import tensorflow as tf x = tf.Variable([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: y = tf.reduce_sum(x ** 2) grad = tape.gradient(y, x) print(grad) # tf.Tensor([2. 4. 6.], shape=(3,), ...)

GradientTape的设计非常明确:你告诉框架“这一段我要记录”,它就在上下文中录下所有涉及tf.Variable的操作。这样不需要给每个张量设置requires_grad,但代价是需要小心管理tape的作用域。如果你在with tf.GradientTape()里调用了不相关的计算,也会被一并记录,影响效率。

JAX 的自动修微分则是函数变换:

import jax import jax.numpy as jnp def loss_fn(x): return jnp.sum(x ** 2) grad_fn = jax.grad(loss_fn) print(grad_fn(jnp.array([1.0, 2.0, 3.0]))) # [2. 4. 6.]

jax.grad接收一个标量输出函数,返回它的梯度函数。你可以继续组合jax.jit(grad_fn)jax.vmap(grad_fn)。这套 API 的好处是组合性和可编译性极强。不过 JAX 的grad默认要求loss_fn是纯函数,也就是说它不能随意读取外部的可变状态,否则梯度结果可能不符合预期。这也解释了为什么 JAX 在训练循环里鼓励把所有参数封装成 PyTree 传入函数。

2.3 模型构建抽象

PyTorch 的模型定义围绕nn.Module展开。你需要继承nn.Module,在__init__里定义子模块,在forward里写前向计算。这非常接近 Python 的面向对象风格。

import torch.nn as nn class MLP(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.fc1 = nn.Linear(in_dim, hidden_dim) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_dim, out_dim) def forward(self, x): return self.fc2(self.relu(self.fc1(x)))

nn.Module会自动收集所有子模块的参数,调用model.parameters()可以直接给优化器。你还可以通过requires_grad = False冻结一部分模块,方便做迁移学习和微调。这也是很多人在网上搜索“PyTorch 冻结部分模型”的原因,因为这套机制确实好用。

TensorFlow 的高层模型 API 是tf.keras.Model。你可以用Sequential快速堆叠,也可以继承Model重写call

import tensorflow as tf class MLP(tf.keras.Model): def __init__(self, hidden_dim, out_dim): super().__init__() self.fc1 = tf.keras.layers.Dense(hidden_dim, activation='relu') self.fc2 = tf.keras.layers.Dense(out_dim) def call(self, x): return self.fc2(self.fc1(x))

Keras 封装度高,容易上手,但在需要精细控制变量创建和共享时,反而比 PyTorch 麻烦。JAX 这边没有官方的高层模型库,常见做法是定义普通 Python 类或纯函数,把权重作为参数字典传入。下面是一个示意性的函数式写法:

import jax.numpy as jnp import flax.linen as nn class MLP(nn.Module): hidden_dim: int out_dim: int @nn.compact def __call__(self, x): x = nn.Dense(self.hidden_dim)(x) x = nn.relu(x) x = nn.Dense(self.out_dim)(x) return x

Flax 是 JAX 生态里常用的模型库,它用nn.Module来定义结构,但参数和模型本身是分离的,需要调用model.init(key, x)来初始化参数。对第一次接触 JAX 的人来说,最难接受的不是没有模型类,而是“参数需要手动传递”这件事。

2.4 训练循环

PyTorch 的经典训练循环是显式 for 循环,一眼能看到发生了什么:

# 伪代码:显式循环 for x, y in dataloader: optimizer.zero_grad() loss = loss_fn(model(x), y) loss.backward() optimizer.step()

这个形式的优点是透明。很多论文里的算法(比如强化学习 TD3)需要频繁更新多个网络、控制随机种子、交错采样和训练,这类需求在显式循环里写起来很舒服。PyTorch Lightning 这类库又在显式循环之上提供了更高封装,如果你想保留控制权又不想写重复模板,可以选它。

TensorFlow 的默认路线是 Keras 的compile+fit

model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) model.fit(train_dataset, epochs=10)

这套 API 接近传统机器学习,代码很短,适合快速验证。但一旦需要自定义损失、梯度裁剪、多塔结构,你仍然需要回到GradientTape手动写循环:

for x, y in train_dataset: with tf.GradientTape() as tape: loss = loss_fn(model(x, training=True), y) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))

Keras 封装并不限制你写底层循环,但很多人会发现自己最终还是在写类似 PyTorch 的代码,只是换成了tf.Variabletape。JAX 的训练循环更强调“训练步函数 + 函数变换”:

@jax.jit def train_step(params, x, y): loss, grads = jax.value_and_grad(loss_fn)(params, x, y) params = jax.tree_map(lambda p, g: p - lr * g, params, grads) return params, loss

这个函数需要你自己更新参数,并且要处理 PyTree 这种结构。刚开始比较麻烦,但一旦适配了这种思路,就可以把jitvmappmap组合到同一个函数里,获得很强的性能优化空间。

2.5 数据加载与预处理

PyTorch 的数据加载是Dataset+DataLoader。你继承Dataset实现__len____getitem__,然后DataLoader负责采样、打乱、多进程读取、自动 batch、collate_fn等。这套 API 的思维是“先定义单个样本怎么取,再自动组织批次”,修改起来很灵活。

TensorFlow 主要用tf.data.Dataset,更接近“数据管道”思维。你从一个数据源开始,然后用链式操作:dataset.map(...).batch(16).prefetch(1),框架会按图执行流水线。它和 Keras 的fit配合得很好,但在写自定义数据预处理时,map里最好用 TensorFlow 算子,直接用 Python 循环或 NumPy 会拖慢速度。

JAX 没有自己的官方 DataLoader。很多人直接把 PyTorch 或 TensorFlow 的数据管道生成 NumPy 数组,再在训练前转成jnp.ndarray。也可以使用tf.data配合 JAX 的input,但需要额外桥接。简单场景下,你可以先用 NumPy 构造数据,再在jax.jit训练函数中传入。这种“缺一个官方加载器”的状态,是 JAX 生态不成熟的表现,也是很多人从 PyTorch 切到 JAX 时最不习惯的地方。

2.6 设备管理与分布式

PyTorch 的设备管理是显式调用model.to(device)device可以是'cuda''cpu'。你要记得把模型和数据都搬到同一设备上。分布式训练有DataParallel和更推荐的DistributedDataParallel。后者需要启动多进程,并在代码里用环境变量初始化进程组。这个设计灵活,但学习曲线并不低。

TensorFlow 的设备管理更偏自动化,GPU 默认可见,关键操作通常会自动分配设备。分布式训练用tf.distribute.Strategy,例如MirroredStrategy,只需要在建立模型和数据集前放入strategy.scope()。你的训练代码可以尽量少改。缺点是当自动策略不生效时,排查起来比较隐蔽。

JAX 的设备管理是一个显式的设备数组概念。jax.devices()返回可用设备列表,数据默认是 CPU 上的普通数组,设备之间的传输由计算触发。你通常不写.cuda(),而是用jax.device_putpmap把计算分布到设备。这种方式更函数式,却也让很多新人在“数据到底在哪一块 GPU 上”这个问题上犯迷糊。

2.7 模型导出与部署

PyTorch 的部署路径较成熟的是 TorchScript,但实际用起来并不轻松。你可以用torch.jit.trace跟踪一个模型,也可以写 TorchScript 脚本,但动态控制流和 Python 语法支持都有限。近年社区也在推torch.export,本质上还是把动态模型变成一张静态计算图。如果你要在服务端部署,通常还要搭配 ONNX、TensorRT 等工具链。

TensorFlow 的部署体系是三家中最完整的。模型可以导出为 SavedModel,然后通过 TensorFlow Serving 加载并对外提供服务;移动端可以用 TFLite,浏览器里可以用 TF.js。Keras 模型导出几乎是一行命令。如果目标是嵌入式、移动端或大规模在线推理,TensorFlow 的历史积累是明显优势。

JAX 的导出主要靠jax2tf,把 JAX 函数转换成 TensorFlow 的 SavedModel 格式。这样做的好处是能借用 TensorFlow 的部署链路;缺点是转换过程有一定限制,而且很多人是希望用 JAX 做研究,部署时才切回 TensorFlow。如果你只是跑实验,可以把导出问题放后面,但如果项目从一开始就考虑上线,这个维度会直接影响框架选型。

3. 环境与安装:API 再强大,装不上也是白搭

很多人一上来就复制安装命令,但最影响后续效率的其实是环境管理方式。我通常在 Windows 或 Linux 上都会先用 Anaconda 创建独立虚拟环境,再安装框架,避免不同项目之间的 Python 版本、CUDA 版本和包依赖互相干扰。

3.1 用 Anaconda 建虚拟环境,先隔离再安装

无论你最终选 PyTorch 还是 TensorFlow,创建一个干净的环境都值得。以 PyTorch 为例,常见的操作是:

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

Python 版本不用追最新,选择一个框架官方已验证过的稳定版本更安全。如果你搜索“PyTorch 环境搭建”或“Anaconda 配置 PyTorch 环境”,大多数教程都会建议这一步。不要直接在 base 环境里安装,因为你的基础环境可能已经有其他项目依赖,版本冲突后很难清理。

虚拟环境隔离还有一个好处:如果安装失败,你不需要卸载一堆包,直接删掉环境重来即可。这个操作成本很低,却能节省大量排查时间。

3.2 PyTorch 安装:先确定 CUDA 版本再选 pip/conda 源

PyTorch 安装最常见的坑是“CPU 版本倒是装上去了,GPU 版本却总是不对”。GPU 版安装前,先确认你的 NVIDIA 显卡驱动支持什么 CUDA 版本。可以在终端执行nvidia-smi查看右上角的 CUDA 版本。然后去 PyTorch 官网选择对应命令,例如:

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

这里cu121表示 CUDA 12.1 的预编译版本。不要只看这个数字,版本号会变,最终要以官网生成命令为准。如果你是 NVIDIA 5000 系显卡,比如搜索里常见的“5060 安装 PyTorch”,要注意太新的 GPU 可能对 CUDA 版本有要求,最好安装较新的 PyTorch 版本,避免不识别。

Jetson 场景也有自己的特殊性。Jetson 的 JetPack 版本不同,预编译的 PyTorch 版本也不一样。搜索“Jetson JetPack 6.2.2 安装什么版本 PyTorch”这类问题时,不能直接用 x86 环境下的安装命令,而是要去 NVIDIA 官方开发者页面找对应设备平台的.whl文件。这类设备安装失败,很多时候不是命令错了,而是平台没选对。

如果你看到 PyTorch 2.6 之后某些老代码加载模型时报错,很可能是torch.loadweights_only默认值发生了变化。这是升级后比较容易踩的兼容性问题,处理方式不是盲目关掉weights_only,而是重新审视模型文件的来源和可信度。

3.3 TensorFlow 2.18 的版本匹配问题

TensorFlow 安装相对简单,基础命令是:

pip install tensorflow

2.x 之后已经不带tensorflow-gpu这个单独包名了,安装tensorflow会在环境具备 GPU 驱动时自动支持 GPU。不过很多人搜索“TensorFlow 2.18 安装”,是因为版本更新后,Python 版本、CUDA 和 cuDNN 的匹配要求发生了变化。

我建议安装前创建干净的虚拟环境,然后先装好对应版本的 NumPy,再安装 TensorFlow,避免依赖解析冲突。如果你在 Windows 上安装,很多时候会遇到 DLL 加载失败或缺少msvcp140.dll,这属于系统运行库问题,需要先安装 Microsoft Visual C++ Redistributable。Linux 上则更多见 CUDA 库路径不对。最好先确认你的环境里有没有预期版本的 CUDA 和 cuDNN,没有的话,也可以用 pip 装带 GPU 支持的 TensorFlow,它会自动拉取一些依赖库,但有时仍需要系统级驱动。

3.4 JAX 的 CPU 与 GPU 安装

JAX 安装是最容易让人困惑的。CPU 版本可以简单执行:

pip install jax jaxlib

GPU 版本则要看你的 CUDA 版本和系统平台,例如pip install jax[cuda12]。由于 JAX 更新频率高,直接给一条命令很可能过时,最稳妥的方式是打开官方文档,根据你用的安装平台选择命令。JAX 没有像 PyTorch 那样有一个统一的官网选版页面,所以经常出现“CPU 能装上,GPU 不生效”的问题。检查方式是在 Python 里执行jax.devices(),如果输出包含CudaDevice说明 GPU 可用,如果只有CpuDevice,则说明 jaxlib 或 CUDA 库有问题。

3.5 安装失败排查顺序

不管哪个框架,安装失败都可以按这个顺序排查:

  1. 先看报错阶段:是命令找不到、依赖冲突、下载超时,还是 import 时报错。
  2. 再看 Python 版本:是否在框架支持范围内。
  3. 检查虚拟环境:是否激活了正确的环境。
  4. 看 pip/conda 源:是否因为网络源不稳定导致安装不完整。
  5. 检查显卡驱动:nvidia-smi是否正常,驱动是否支持目标 CUDA。
  6. 看包版本:版本号是否与框架匹配。
  7. 最后看硬件平台:是 x86 还是 ARM/Jetson,有没有使用特殊安装源。

这套排查链路能覆盖绝大多数情况。如果环境反复损坏,建议直接新建虚拟环境,而不是在旧环境里继续卸载重装。不要一开始就去改整合包或系统级 Python,很多额外问题都是因为环境被手工弄得过于复杂。

注意:不要一上来就把批量数和并发数拉满,先用一条样例确认输入、输出和日志都正常。

4. 选型建议:别只看生态,还要看你的训练循环写了多少

4.1 选 PyTorch 的场景

如果你要做研究、快速验证算法、复现论文,PyTorch 是大多数情况下的首选。它的动态图和显式训练循环让自定义代码路径变得很自然,社区几乎已经把 PyTorch 当成了默认语言。你在网上搜“PyTorch 实现 Transformer”“CycleGAN 代码”“猫狗分类”“手写数字识别”,甚至“目标检测”,会发现绝大多数开源项目默认给出 PyTorch 版本。强化学习里的 TD3、DQN 这类频繁修改网络结构的算法,用 PyTorch 的显式循环写起来也更顺手。

PyTorch 还比较适合一个人掌握全局的小项目。你可以只依靠torch.nntorchvision快速构建模型,再用DataLoader管理数据。无论你是做图像、文本还是多模态,资料都要比其他框架多很多。

4.2 选 TensorFlow 的场景

选择 TensorFlow 的理由通常不在“写起来多舒服”,而在于生产链路。如果你的项目需要在服务端高并发推理、部署到移动端、转成 TFLite,或者团队已经有一套基于 TF Serving 的基础设施,那 TensorFlow 是不错的选择。Keras 的compile+fit也能帮助更快速地做表格数据或传统图像分类实验。

不过要注意,TensorFlow 的 API 历史包袱比较重,很多教程还是旧版 Session 写法,新手很容易被陈旧资料带偏。如果你决定用 TensorFlow,请直接看 TensorFlow 2.x 和 Keras 的官方文档,并尽量使用高层接口。一旦需要深度自定义训练,成本会比 PyTorch 高一些,你需要同时理解 Keras 封装和底层GradientTape机制。

4.3 选 JAX 的场景

JAX 适合愿意接受函数式风格、追求高性能和可扩展性的人。如果你研究的是大规模模型、TPU 训练、自动向量化,或者想尝试一种能同时写出简洁数学逻辑和高效执行代码的方式,JAX 会带来惊喜。通过jax.vmap可以自动把 batch 维度加上去,不用到处改循环;通过jax.jit可以把计算逻辑编译成高效的 XLA 内核。

但如果你只是为了尽快完成一篇论文、跑通一个常规项目,JAX 的学习成本可能超过收益。它缺少官方数据加载器和一套默认模型库,选 Flax、Haiku 还是 Equinox 本身就是一个学习负担。JAX 更适合“第二个来学习的框架”,而不是“第一个入门框架”。

4.4 一个四步判断框架

面对团队或个人的框架选择,我建议用下面四步做判断,而不是只看流行度:

  1. 先看部署终点:如果最终要上线移动端或服务端,TensorFlow 的导出链更成熟;如果只是实验脚本,PyTorch 更节省时间。
  2. 再看训练循环复杂度:如果你的算法需要大量自定义控制流,PyTorch 的显式循环更友好;如果模型结构稳定且能接受fit风格,TensorFlow 更高效。
  3. 看社区资源和技术债:你项目中已有的代码、团队擅长什么、能否维护。不要为了新而新,把团队拖入不熟悉的地带。
  4. 最后看硬件和性能需求:如果要用 TPU 集群或大规模并行,JAX 值得试;如果是普通单卡开发,PyTorch 仍然是较低风险的选择。

三套框架未来大概率还会继续互相借鉴:PyTorch 在加强部署能力,TensorFlow 在优化开发体验,JAX 在不断扩展生态。作为使用者,不用把选型看成“信仰之争”。先把七个核心 API 维度的差异化成自己的坐标系,等下一个新框架出现时——无论它是不是深度学习框架——你都能快速定位到它的张量 API、自动微分、模型构建、训练循环、数据管道、设备管理和部署路径,这才是对比的真正价值。下一步最该做的,不是再收藏一篇对比文章,而是拿一个手写数字识别这样的小任务,选一个框架把最小训练循环跑通。环境、报错和 API 手感,跑一遍体会会更深。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询