十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch、TensorFlow、JAX核心API对比与安装指南

PyTorch、TensorFlow、JAX核心API对比与安装指南 深度学习框架的核心 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.grad、jax.jit、jax.vmap这类函数变换来组织代码。这个区别为什么会直接影响 API因为框架要想做到自动求导必须能追踪计算过程。PyTorch 在张量上用requires_grad标记和动态图回溯TensorFlow 用GradientTape记录正向计算过程JAX 则要求你的损失函数是纯函数然后对它做函数变换。你可以用差不多的数学公式写出同一个模型但训练循环和调试方式会很不一样。最直接的感受是PyTorch 像是给每个张量装了一个记录器TensorFlow 像在旁边放了一卷录音带JAX 更像是一个可以把“函数”整体变形成“导函数”的高阶函数工厂。1.2 出身决定了 API 的性格PyTorch 脱胎于 Torch底层用 C 加速上层用 Python 做交互。它从一开始就选择了“让研究人员先跑起来”的路线所以 API 非常贴 Python 习惯Module、optimizer、dataloader这些对象都很直觉。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_gradTrue) 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重写callimport 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, activationrelu) 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 xFlax 是 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 的compilefitmodel.compile(optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy]) model.fit(train_dataset, epochs10)这套 API 接近传统机器学习代码很短适合快速验证。但一旦需要自定义损失、梯度裁剪、多塔结构你仍然需要回到GradientTape手动写循环for x, y in train_dataset: with tf.GradientTape() as tape: loss loss_fn(model(x, trainingTrue), y) grads tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))Keras 封装并不限制你写底层循环但很多人会发现自己最终还是在写类似 PyTorch 的代码只是换成了tf.Variable和tape。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 这种结构。刚开始比较麻烦但一旦适配了这种思路就可以把jit、vmap、pmap组合到同一个函数里获得很强的性能优化空间。2.5 数据加载与预处理PyTorch 的数据加载是DatasetDataLoader。你继承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_put或pmap把计算分布到设备。这种方式更函数式却也让很多新人在“数据到底在哪一块 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 python3.10 -y conda activate torchPython 版本不用追最新选择一个框架官方已验证过的稳定版本更安全。如果你搜索“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.load的weights_only默认值发生了变化。这是升级后比较容易踩的兼容性问题处理方式不是盲目关掉weights_only而是重新审视模型文件的来源和可信度。3.3 TensorFlow 2.18 的版本匹配问题TensorFlow 安装相对简单基础命令是pip install tensorflow2.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 jaxlibGPU 版本则要看你的 CUDA 版本和系统平台例如pip install jax[cuda12]。由于 JAX 更新频率高直接给一条命令很可能过时最稳妥的方式是打开官方文档根据你用的安装平台选择命令。JAX 没有像 PyTorch 那样有一个统一的官网选版页面所以经常出现“CPU 能装上GPU 不生效”的问题。检查方式是在 Python 里执行jax.devices()如果输出包含CudaDevice说明 GPU 可用如果只有CpuDevice则说明 jaxlib 或 CUDA 库有问题。3.5 安装失败排查顺序不管哪个框架安装失败都可以按这个顺序排查先看报错阶段是命令找不到、依赖冲突、下载超时还是 import 时报错。再看 Python 版本是否在框架支持范围内。检查虚拟环境是否激活了正确的环境。看 pip/conda 源是否因为网络源不稳定导致安装不完整。检查显卡驱动nvidia-smi是否正常驱动是否支持目标 CUDA。看包版本版本号是否与框架匹配。最后看硬件平台是 x86 还是 ARM/Jetson有没有使用特殊安装源。这套排查链路能覆盖绝大多数情况。如果环境反复损坏建议直接新建虚拟环境而不是在旧环境里继续卸载重装。不要一开始就去改整合包或系统级 Python很多额外问题都是因为环境被手工弄得过于复杂。注意不要一上来就把批量数和并发数拉满先用一条样例确认输入、输出和日志都正常。4. 选型建议别只看生态还要看你的训练循环写了多少4.1 选 PyTorch 的场景如果你要做研究、快速验证算法、复现论文PyTorch 是大多数情况下的首选。它的动态图和显式训练循环让自定义代码路径变得很自然社区几乎已经把 PyTorch 当成了默认语言。你在网上搜“PyTorch 实现 Transformer”“CycleGAN 代码”“猫狗分类”“手写数字识别”甚至“目标检测”会发现绝大多数开源项目默认给出 PyTorch 版本。强化学习里的 TD3、DQN 这类频繁修改网络结构的算法用 PyTorch 的显式循环写起来也更顺手。PyTorch 还比较适合一个人掌握全局的小项目。你可以只依靠torch.nn和torchvision快速构建模型再用DataLoader管理数据。无论你是做图像、文本还是多模态资料都要比其他框架多很多。4.2 选 TensorFlow 的场景选择 TensorFlow 的理由通常不在“写起来多舒服”而在于生产链路。如果你的项目需要在服务端高并发推理、部署到移动端、转成 TFLite或者团队已经有一套基于 TF Serving 的基础设施那 TensorFlow 是不错的选择。Keras 的compilefit也能帮助更快速地做表格数据或传统图像分类实验。不过要注意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 一个四步判断框架面对团队或个人的框架选择我建议用下面四步做判断而不是只看流行度先看部署终点如果最终要上线移动端或服务端TensorFlow 的导出链更成熟如果只是实验脚本PyTorch 更节省时间。再看训练循环复杂度如果你的算法需要大量自定义控制流PyTorch 的显式循环更友好如果模型结构稳定且能接受fit风格TensorFlow 更高效。看社区资源和技术债你项目中已有的代码、团队擅长什么、能否维护。不要为了新而新把团队拖入不熟悉的地带。最后看硬件和性能需求如果要用 TPU 集群或大规模并行JAX 值得试如果是普通单卡开发PyTorch 仍然是较低风险的选择。三套框架未来大概率还会继续互相借鉴PyTorch 在加强部署能力TensorFlow 在优化开发体验JAX 在不断扩展生态。作为使用者不用把选型看成“信仰之争”。先把七个核心 API 维度的差异化成自己的坐标系等下一个新框架出现时——无论它是不是深度学习框架——你都能快速定位到它的张量 API、自动微分、模型构建、训练循环、数据管道、设备管理和部署路径这才是对比的真正价值。下一步最该做的不是再收藏一篇对比文章而是拿一个手写数字识别这样的小任务选一个框架把最小训练循环跑通。环境、报错和 API 手感跑一遍体会会更深。
返回列表