ARTICLE DETAIL

资讯详情

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

深度学习框架PyTorch、TensorFlow、JAX核心API对比与迁移指南

深度学习框架PyTorch、TensorFlow、JAX核心API对比与迁移指南 深度学习框架选型核心不是看哪个框架热搜多而是看核心 API 的用法是否匹配你的项目阶段。PyTorch、TensorFlow、JAX 这三个框架我从安装、建模、训练到部署都跑过最直接的感受是PyTorch 把动态图做成了默认体验TensorFlow 把工程链路做得很完整JAX 则把函数式变换做到了极致。下面按实际选型和迁移的顺序把七个常见框架的核心 API 逐个拆解重点放在 PyTorch 对比 TensorFlow 对比 JAX 的建模、训练、数据加载和序列化环节顺便说清楚什么时候选谁以及换框架时到底要重写哪些代码。1. 先搞懂“核心API”到底在比什么很多人看框架对比第一步就去看训练速度、显存占用、模型库数量。这些指标当然重要但真正决定迁移成本的是核心 API 的设计方式。同一个多层感知机在 PyTorch 里继承一个类写 forward在 TensorFlow 里调用 Keras 的 compile 和 fit在 JAX 里定义纯函数再用 grad 变换代码组织方式完全不同。所以对比核心 API不能只比谁的文档好看要比下面这几层。1.1 模型定义接口的差异决定代码组织方式模型定义是框架和用户接触最深的一层。PyTorch 用nn.Module把模型封装成类用户通过__init__声明子模块在forward里写张量运算。由于动态图机制forward 内部可以写 Python 原生的 if、for、print甚至临时改逻辑。这对调试和算法研究非常友好。TensorFlow 2.x 的主推方式是 Keras。用Sequential或函数式 API 把层串起来代码非常短。但短也意味着控制力相对收敛一些奇怪的动态分支写起来更别扭。虽然 Keras 也支持子类化但官方在很多场景下推荐的是函数式 API因为函数式更方便序列化和图编译。JAX 干脆没有内置模型类。它提供的是jax.numpy、jax.grad、jax.jit这类函数变换工具模型结构需要靠 Flax 或 Haiku 这类第三方库组织。初学者第一次从 PyTorch 切到 JAX最不适应的不是 API 少而是“模型”这个概念变轻了本质上就是参数和纯函数。1.2 自动求导、数据管道和部署接口才是迁移成本大头模型定义之外自动求导方式决定了训练代码能不能随心所欲地改造。PyTorch 的 autograd 在每次 forward 后自动建图backward()直接回传梯度写自定义 loss 和梯度裁剪都很顺手。TensorFlow 用tf.GradientTape记录梯度控制流更显式但和 Keras 高层 fit 混用时你会发现自己反而需要跳出封装去写底层逻辑。JAX 则要求你先把损失函数定义成纯函数再用jax.grad生成梯度函数所有状态都是显式传入和返回。数据管道也是重头戏。PyTorch 的DataLoader基于Dataset和Sampler适合自定义数据集配合num_workers就能起多进程加载。TensorFlow 的tf.data更偏向流式和图化复杂预处理写起来和普通 Python 风格不太一样。JAX 官方没有绑定 DataLoader一般用tf.data或 NumPy 切片再手动 batch数据 shuffle、并行、重复的规则都要自己理清楚。部署接口更不能忽略。PyTorch 有 TorchScript、Torch.compile、ONNX 导出最近几版还在做新的导出链路TensorFlow 的核心产物是 SavedModel配合 TensorFlow Serving 可以很快搭出服务JAX 因为有 jit 和 pmap理论上很优雅但落地服务时通常需要先导出到 SavedModel 或 ONNX 这类中间格式链路比前两者要绕一些。2. PyTorch动态图优先API风格最接近原生PythonPyTorch 是目前研究场景里最主流的选择。它的核心 API 不是一个大一统接口而是 Tensor、autograd、nn.Module、DataLoader这几个模块组合在一起。只要你能接受“训练循环自己写”大部分事情做起来都比较直接。2.1 nn.Module、autograd 和 DataLoader 的配合方式PyTorch 的建模路径很清楚定义类继承nn.Module在__init__里实例化需要的层在forward里写计算逻辑。比如定义一个两层 MLP代码就是类 forward中间想加打印、分支、调试断点都可以。import torch 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)))训练时DataLoader负责把数据集按 batch 吐出来。batch_size、shuffle、num_workers是最常调的三个参数。很多人上来就把num_workers调得很大结果 Windows 环境下频繁报错或者内存暴涨。我的建议是小任务先用num_workers0或1跑通再去压多进程加载。autograd 会在 forward 执行时自动记录每个张量的操作然后loss.backward()把梯度传到每个叶子张量。这个机制让自定义 loss 很方便但也要注意如果你在 forward 里直接修改requires_gradTrue的张量或者用了原地操作梯度链路可能断掉。遇到梯度为 None 时先查是不是 forward 里用了不合适的 index 赋值。2.2 显式训练循环的优势调试直观控制力强PyTorch 最常见的训练循环是自己写for x, y in dataloader: opt.zero_grad() out model(x) loss loss_fn(out, y) loss.backward() opt.step()这段代码看着简单但它把训练的真实步骤完全暴露出来了。你想在 backward 之前做梯度裁剪直接加一行想打印某个中间层输出直接 forward 里 print想冻结部分参数遍历model.parameters()改requires_grad就行。缺点是项目一多循环代码会重复。于是出现了 PyTorch Lightning、Accelerate 这类封装。它们不是新框架而是把重复的循环、日志、checkpoint、分布式逻辑收起来。使用时要清楚方便性提升的同时你也部分放弃了直接控制。出问题先看封装库的版本是否和 PyTorch 匹配很多人报一堆莫名其妙的错最后发现是 Lightning 版本太老。2.3 需要注意的版本变化和序列化问题PyTorch 的 API 变动不算频繁但大版本升级时总有意外。比如 PyTorch 2.6 开始对torch.load的weights_only参数默认值做了调整升级后加载旧版 checkpoint 可能会因为安全限制报错。我在验证旧模型时遇到过类似情况处理方式不是直接改源码而是先确认当前 PyTorch 版本和 checkpoint 来自哪个版本再有针对性地设置加载参数。序列化是很多项目最后才遇到的问题。torch.save保存的不只是模型参数还可能是 Python 对象的引用路径。如果你把模型类和训练脚本放在不同目录换机器加载时可能出现找不到自定义类的错误。稳妥做法是只保存state_dict再在目标环境中重新实例化模型结构。3. TensorFlow从Keras到Graph接口分工明显TensorFlow 过去因为 1.x 和 2.x 的割裂劝退了不少人。2.x 之后官方把 Keras 定为高层面向前端用户日常开发基本都是 Keras。但 TensorFlow 的真正优势并不在建模而在生产链路tf.data、SavedModel、Serving、TFX 这些工具把训练和部署串得很完整。3.1 Keras高层接口把训练流程收得很紧Keras 的核心 API 是Sequential、函数式 API 和子类化三种建模方式。最常见的机器学习任务可以直接写import tensorflow as tf model tf.keras.Sequential([ tf.keras.layers.Dense(hidden_dim, activationrelu, input_shape(in_dim,)), tf.keras.layers.Dense(out_dim) ]) model.compile(optimizeradam, losstf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue)) model.fit(train_dataset, epochs10)compile和fit把训练过程封装得很干净。你要做的只是指定优化器、损失函数、评估指标和数据。对于标准分类、回归任务这套接口比手写循环省事很多。但正因为封装遇到不标准的需求时会比较难受。比如想在每个 step 里做自定义梯度修改就得放弃fit改用tf.GradientTape手写训练循环。Keras 的子类化也能处理复杂模型但子类模型的序列化和部署灵活性比函数式模型差一些。我的建议是能用函数式 API 就别轻易上子类化除非模型结构确实需要动态控制流。3.2 tf.function、tf.data和SavedModel是生产化关键TensorFlow 和 PyTorch 最大的区别是它始终把“图”作为生产基础。tf.function可以把普通 Python 函数编译成 TensorFlow 图配合tf.data构建输入管道再用SavedModel统一导出。这样训练端到部署端的格式很一致。tf.data的接口和普通 Python 迭代器不一样。它习惯用map、batch、shuffle、repeat、prefetch这一套链式调用。切到 TensorFlow 时最容易踩的坑是在map里写纯 Python 逻辑或者用了外部全局变量导致图编译失败。尽量把预处理写成tf.Tensor操作或者用tf.py_function包一层但后者会影响性能。SavedModel 是 TensorFlow 的通用导出格式。它把模型结构、权重、签名一起打包部署时用 TensorFlow Serving 加载。相比裸权重文件SavedModel 更适合作 API 服务因为可以通过签名定义输入输出字段。做生产项目时不要只在训练脚本里保存 h5最好把 Serving 的请求和响应格式也一起规划。3.3 安装和CUDA版本匹配容易踩坑每次 TensorFlow 大版本发布最先被搜索的往往不是新功能而是怎么安装。TensorFlow 2.18 这种版本对 Python、CUDA、cuDNN 的版本匹配非常敏感。我一般建议先建干净的虚拟环境再按官方 pip 安装文档选择命令而不是直接在全局环境里pip install tensorflow。GPU 版更麻烦。很多人的报错并不是模型代码问题而是 CUDA 动态库找不到或者 cuDNN 和 TensorFlow 版本不匹配。处理顺序一般是先确认nvidia-smi里的驱动版本再确认能支持的 CUDA 版本最后按官方对应关系安装 TensorFlow。不要只看驱动版本就装最新 CUDA驱动、CUDA、TensorFlow 三层要同时成立。如果你在 Jetson 这类嵌入式设备上跑情况又不一样。JetPack 系统往往已经有定制好的 PyTorch 或 TensorFlow wheel不能直接拿 x86 的安装命令硬套。比如 JetPack 6.2.2 到底对应哪个 PyTorch 版本要看官方 release 说明而不是猜。这类问题先去厂商论坛或官方文档确认比换源重装更有效。4. JAX函数式变换是灵魂所有API都围绕“纯函数”JAX 不是一个传统意义上的深度学习框架。它更像一套可微分的 NumPy 加函数变换工具。核心 API 不是Model而是jax.numpy、jax.grad、jax.jit、jax.vmap、jax.pmap。习惯了 PyTorch 的类和对象风格后切到 JAX 会需要一点思维转换。4.1 jax.numpy、grad、jit、vmap、pmap的关系JAX 的jax.numpy尽量兼容 NumPy 的 API但它有两个硬约束数组不可变函数必须纯净。不可变意味着x[idx] value这类原地修改是不行的你得用x.at[idx].set(value)。纯净意味着函数不能修改全局变量不能依赖随机状态随机数需要自己显式传 key。jax.grad是对纯函数求梯度的核心 API。它要求目标函数的输入输出都是数组或数组结构返回值通常是一个标量损失。梯度函数和原函数的签名略有不同返回的是第一个参数对应的梯度。jax.jit会把函数编译成 XLA 算子加速效果明显但编译期间不能有 Python 副作用。jax.vmap用来把单样本函数向量化成 batch 版本不用手动写 batch 维逻辑。jax.pmap则把计算复制到多设备上并行执行。它们看起来是几个小函数但组合起来能覆盖很多分布式场景。需要注意的是这些函数叠加使用时输入输出结构必须一致否则报错信息会非常晦涩。4.2 Flax和Haiku怎么组织模型JAX 官方不提供厚重的模型类但生态里有两个主流选择Flax 和 Haiku。Flax 的linen模块提供类似 PyTorch 的类定义方式Haiku 是 DeepMind 风格用hk.Module加hk.transform。下面用 Flax 写一个 MLPimport jax.numpy as jnp from flax import 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和 PyTorch 不同的是Flax 的层通常在__call__中即时实例化参数由框架管理需要初始化后才能拿到参数结构。训练时要把参数和 batch 数据一起传入纯函数梯度更新也要手动完成。这个流程第一次跑通时会觉得繁琐但跑顺手后你会对“参数在哪儿、状态在哪儿”特别清楚。4.3 调试思路和PyTorch、TensorFlow完全不同JAX 的调试是很多人卡住的地方。因为jit和纯函数限制你不能在编译函数里随便print中间张量。常规做法是先在纯 NumPy 模式下跑通函数再把jit加回去。确实需要打印时可以用jax.debug.print这类显式打印接口但不能依赖普通print。还有一个高频坑随机数。PyTorch 里直接torch.randn很方便JAX 里则必须创建 key再用jax.random.split拆 key 传入函数。如果不小心在函数外面复用了同一个 key不同 batch 的随机数可能一模一样。这个问题看起来像“模型输出没变化”实际是随机种子没管好。5. 同一个任务三个框架的代码差异有多大代码对比是最直观的。这里用一个多层感知机、一个注意力模块、一个训练循环来展示差异。你会发现模型定义只是表象真正的差异在训练和数据处理上。5.1 多层感知机模型定义对比PyTorch 是类继承nn.ModuleTensorFlow 是 Keras 层堆叠JAX 是 Flax 的nn.Module配合nn.compact。三层写法从上文已经能看到。如果只看这几段代码似乎没有天壤之别。但放到完整训练上下文里差异就放大了。PyTorch 的模型实例是可调用的 Python 对象参数在model.parameters()里优化器直接吃这个生成器。Keras 的模型自带compile和fit训练逻辑被收纳进高层接口。Flax 的模型对象更像一个描述结构真正的状态是params字典需要你手动传进函数。三种写法对应了三种工程哲学手动控制、自动封装、纯函数传递。5.2 Transformer/注意力模块核心API对比注意力模块是现在调用频率最高的 API 之一。三个框架都内置了多头注意力但接口差异明显框架核心API输入形状特点输出内容PyTorchtorch.nn.MultiheadAttention输入形状是(seq_len, batch, embed)批次维在中间返回输出和注意力权重TensorFlowtf.keras.layers.MultiHeadAttention输入形状是(batch, seq, embed)批次维在最前返回注意力输出可配置是否返回权重JAX/Flaxflax.linen.MultiHeadAttention输入形状是(batch, seq, features)批次维在最前返回输出需要自己决定是否取中间变量刚切框架的人最容易在这里踩坑PyTorch 和 TensorFlow/JAX 的 batch 维位置不一样。同样的 Transformer 代码从 PyTorch 迁到 TensorFlow经常要把permute或transpose调整一遍否则维度不匹配报错。很多报错看起来像模型结构写错了实际是输入张量的维度约定没对齐。5.3 训练循环fit与手写循环的取舍PyTorch 的默认训练循环是显式 for 循环for x, y in dataloader: opt.zero_grad() out model(x) loss loss_fn(out, y) loss.backward() opt.step()TensorFlow 的默认训练是 Keras fitmodel.compile(optimizeradam, losssparse_categorical_crossentropy) model.fit(train_dataset, epochs10)JAX 则要自己写参数更新函数def loss_fn(params, x, y): out model.apply(params, x) return jnp.mean((out - y) ** 2) grad_fn jax.grad(loss_fn) grads grad_fn(params, x, y) params jax.tree_util.tree_map(lambda p, g: p - lr * g, params, grads)放到一起看Keras 最短PyTorch 最直接JAX 最显式。选型时不要只看哪个代码更短还要看项目是否需要自定义训练逻辑。Keras 的高层封装适合标准任务PyTorch 的手写循环适合快速实验JAX 的纯函数适合算法级定制和分布式扩展。6. 七个框架核心API速览与选型建议除了 PyTorch、TensorFlow、JAX完整对比还需要把 Keras、PaddlePaddle、MindSpore、MXNet 放进来。严格说Keras 是 TensorFlow 的前端 APIPaddlePaddle 和 MindSpore 是独立的国产框架MXNet 更多存在于历史项目里。把它们放在一起看能更清楚不同框架在 API 设计上的取舍。6.1 七个框架在核心API上的定位差异框架核心API模型定义方式训练方式当前适合场景PyTorchtorch.nn、autograd、DataLoader类继承nn.Module动态图手写循环为主也可用 Lightning研究、原型验证、中小型训练TensorFlowtf.keras、tf.function、tf.dataSequential、函数式、子类化compilefit或自定义循环生产链路、大规模服务化JAXjax.numpy、grad、jit、vmap、pmapFlax、Haiku 模型库纯函数式手写训练大规模并行、科学计算、算法研究Keraskeras.layers、compile、fit序列式、函数式、子类化高层训练接口快速建模、多后端前端PaddlePaddlepaddle.nn、paddle.static类继承paddle.nn.Layer动态图为主高层 API 或手写循环国内产业落地、飞桨生态MindSporemindspore.nn、mindspore.Cell类继承Cell支持图和动静态高阶 API 或函数式训练昇腾硬件、端边云协同MXNetmxnet.gluon、mxnet.symbolGluon 模块化或 Symbol 静态图gluon.Trainer历史项目、课程遗留代码这个表不是为了排先后而是说明各框架的核心 API 对应的是不同视角。PyTorch 和 MXNet 的 Gluon 都偏动态TensorFlow 和 Keras 偏工程封装JAX 偏函数变换PaddlePaddle 和 MindSpore 则和自家硬件、平台绑定更紧。6.2 选型不是看谁火而是看项目阶段我的经验是如果做研究、复现论文、快速验证想法优先 PyTorch因为它的社区和预训练模型资源最全。如果做线上 API 服务并且团队已经有标准化部署体系TensorFlow 的 SavedModel 和 Serving 链路更省心。如果你要做超大模型并行或者可微分编程这类偏底层的事JAX 的函数式 API 值得花时间。PaddlePaddle 和 MindSpore 的选择更多要看部署环境。公司或实验室如果已经有对应硬件平台比如昇腾设备用 MindSpore 会比强行装 PyTorch 再适配更顺。MXNet 除非维护老项目否则不建议新项目再入它的生态活跃度已经远不如前三强。选型时还要看团队熟悉度。框架可以换但团队的学习成本是实际的。如果大家已经习惯 PyTorch 的调试方式硬换 TensorFlow 只会让开发变慢。技术选型不是选一个“最好”的而是选一个当前团队、硬件、业务约束下最顺的。7. 环境准备和常见坑从安装到第一个Demo最后聊环境。很多人不是被模型代码难住的而是被安装、CUDA、依赖版本搞到崩溃。下面按 PyTorch、TensorFlow、JAX 分别梳理再给一个通用的验证顺序。7.1 PyTorch和TensorFlow的最小安装方案PyTorch 的安装建议先用 conda 建独立环境再去官方安装页选择操作系统、包管理器和 CUDA 版本复制对应的命令。比如conda create -n dl python3.10 conda activate dl pip install torch --index-url https://download.pytorch.org/whl/cu121命令里的 CUDA 版本要和你的环境对应不要无脑复制。如果只是 CPU 学习安装 CPU 版就够没必要先折腾 GPU。GPU 跑不起来先看nvidia-smi是否正确输出再确认 CUDA 版本和 PyTorch wheel 是否匹配。TensorFlow 的安装也差不多的思路。TensorFlow 2.18 安装时先确认当前 Python 版本官方支持到多少再按官方命令安装。GPU 版最容易遇到的是cudart64_*.dll或libcudnn.so找不到这通常不是框架问题而是 CUDA 和 cuDNN 没配对。处理方式不是重装 TensorFlow而是把 CUDA 相关依赖对齐。7.2 JAX安装时最容易忽略的CUDA匹配问题JAX 安装相对特殊因为它要绑定 CUDA 运行时。常见命令是pip install jax[cuda12]这类索引包具体看官方文档。装错 CUDA 版本时运行时报错可能很晚才出现比如一开始 import 正常但执行jax.devices()或跑一个jnp.dot时才崩溃。我踩过的坑是在已经装有系统 CUDA 的机器上直接用pip install jax装了 CPU 版结果所有 GPU 加速都不生效。后来必须显式安装带 CUDA 的版本并把LD_LIBRARY_PATH指向正确位置。JAX 的 GPU 是否可用用jax.devices()验证比看安装返回信息靠谱得多。7.3 跑第一个Demo时建议按什么顺序验证不管用哪个框架我建议第一个 Demo 不要直接跑完整训练而是按下面的顺序验证确认框架能正常导入import torch、import tensorflow、import jax不报错。确认设备可用分别用torch.cuda.is_available()、tf.config.list_physical_devices(GPU)、jax.devices()检查。跑一个极小矩阵乘确认基本运算正常。跑一个单 batch 的前向和反向确认自动求导链路正常。最后才跑完整训练循环。低配置机器也能跑但要把 batch size、分辨率、并发数降下来不要一上来就开最大参数。遇到报错先看现象是在 import 阶段、数据加载阶段还是训练阶段再去看日志。很多问题看起来是模型代码不对最后发现是路径、权限、依赖版本或输入格式的问题。先确认这些再改参数效率会高很多。如果你正在选型我的建议很直接想快速出研究结果和做实验PyTorch 最稳想往生产服务上放TensorFlow 的工程链路值得认真看想研究大规模并行或新算法JAX 值得花时间适应。其他几个框架更多是看生态和硬件绑定而不是看谁好看。真正换框架时最需要重写的不是模型定义本身而是数据管道、训练循环和序列化方式。先把这个意识建立起来选哪个框架都会少踩很多坑。
返回列表