深度学习框架API对比:PyTorch、TensorFlow与JAX等七大框架迁移指南
2026/8/31 9:11:59 网站建设 项目流程

如果你在小团队里既要做算法实验,又要接线上推理,大概率会被几个框架的 API 差异折磨过:PyTorch 里写惯了model(x)loss.backward(),换到 TensorFlow 发现自动微分用的是tf.GradientTape,再换到 JAX 又变成了jax.grad(loss_fn)(params, x),连参数都要自己端在手里。更别提维度顺序、模型保存格式、随机种子这些细节,几乎每个框架都有自己的“个性”。

这篇文章准备把 7 个深度学习框架拉到一起,从张量创建、自动微分、模型构建、训练循环、模型保存五个维度做一次核心 API 对比。重点放在 PyTorch、TensorFlow、JAX 三巨头,同时兼顾 Keras、PaddlePaddle、MindSpore、MXNet 的差异化设计。读完你得到的不是一个 API 清单,而是一张“API 迁移地图”:以后无论切到哪个框架,都能快速定位对应的接口,并避开那些最容易踩的坑。

这里先给一个明确判断:七个框架的 API 差异,本质上不是命名风格不同,而是编程范式不同。PyTorch 是命令式动态图,TensorFlow 是声明式图加高层封装,JAX 是函数式数值计算。理解了一个框架的范式,再看另一个框架的 API,就会觉得“它只是换了一种表达方式”,而不是“又要重新学一遍”。

1. 为什么深度学习框架的 API 差异值得认真对比

很多开发者对框架的选择是“跟风”的:论文用什么我就用什么,公司技术栈是什么我就用什么。真正到了切换框架的时候才发现,把一个训练好的模型从 PyTorch 迁到 TensorFlow,不只是把import torch改成import tensorflow as tf这么简单。

最直接的痛点主要有三个。

第一,自动微分的调用方式完全不同。PyTorch 是“构建计算图 + 调用 backward”,TensorFlow 是“用 GradientTape 记录前向过程再反向求解”,JAX 则是“把损失函数作为纯函数传给 grad”。如果没理解这三种机制的区别,代码报错时很难定位是数学问题还是 API 使用问题。

第二,模型参数的管理方式不同。PyTorch 用nn.Module自动跟踪参数,Keras 用LayerModel跟踪参数,JAX 核心库没有“模型”概念,参数必须显式放在函数签名里,常用FlaxHaikuEquinox这类生态库来辅助。这个差异会直接影响你的训练循环怎么写。

第三,生态和部署链路不同。PyTorch 在科研社区最活跃,TensorFlow 在成熟的企业系统里存量巨大,JAX 在大模型训练和高性能科学计算上增长很快。选框架从来不是“谁更好”的问题,而是“谁更适合你当前的场景”的问题。

所以这篇文章不打算评价哪个框架“最强”,而是想把它们的核心 API 放到同一张表里,帮你看清楚映射关系。对技术读者来说,掌握跨框架抽象能力,比死记单一框架更能应对项目变化。

2. 七个框架的定位与生态:一张表看懂

在进入代码之前,先建立整体认知。七个框架分别是 PyTorch、TensorFlow、JAX、Keras、PaddlePaddle、MindSpore、MXNet。

框架维护方核心编程范式主要接口入口典型场景
PyTorchLinux Foundation / Meta命令式动态图torch.nntorch.optimtorch.autograd科研实验、快速原型、TorchServe 部署
TensorFlowGoogle声明式图 + 动态执行tf.kerastf.datatf.function生产环境、移动端、成熟 MLOps 链路
JAXGoogle Research函数式 + XLA 编译jax.numpyjax.gradjax.jitjax.vmap高性能科学计算、大规模并行训练
KerasGoogle + 社区高层声明式keras.Modelkeras.layers快速建模、多后端迁移
PaddlePaddle百度动静态统一paddle.nn.Layerpaddle.to_static本地生态、全流程工业平台
MindSpore华为动静态统一 / 图编译mindspore.nn.Cellmindspore.trainAI 与科学计算融合、特定硬件平台
MXNetApache动态 + 混合编程mxnet.ndmxnet.gluon旧项目维护、教学

注意,Keras 更准确的定位是“高层模型 API”,它从 3.0 开始支持 PyTorch、TensorFlow、JAX 等多个后端。把它放进这个名单,是因为很多 TensorFlow 用户实际接触的是 Keras API,而不是底层图 API。

从整个行业趋势看,PyTorch 在学术论文和开源模型复现中的占比已经明显领先,TensorFlow 在企业旧系统和特定部署场景中仍然有大量存量,JAX 则凭借函数式转换和 XLA 编译,在需要大算力、大规模并行的场景中越来越受关注。PaddlePaddle 和 MindSpore 更多与各自的软硬件生态绑定,MXNet 虽然还在 Apache 旗下,但更新节奏已经放慢,新项目不太建议选择。

3. 核心张量 API 对比:Tensor / Tensor / Array

所有深度学习框架的底层都是多维数组运算。PyTorch 叫torch.Tensor,TensorFlow 叫tf.Tensor,JAX 叫jax.Array,Paddle 叫paddle.Tensor,MindSpore 叫mindspore.Tensor,MXNet 叫mxnet.ndarray.NDArray。虽然名字不同,但它们都要处理三个核心问题:数据类型 dtype、形状 shape、设备 device。

新手最容易忽略的是设备语义。PyTorch 里你能直接看到tensor.device是 CPU 还是 CUDA,TensorFlow 也有类似概念,JAX 则默认把数组看作“不关心设备”的抽象值,统一由jax.jitjax.device_put来管理。这意味着 JAX 代码看起来更干净,但如果你不熟悉它的异步调度,反而可能发现显存占用异常。

下面用同一段“创建 3×4 随机矩阵并计算矩阵乘法”演示三个框架的写法。

# PyTorch import torch device = 'cuda' if torch.cuda.is_available() else 'cpu' x = torch.randn(3, 4, dtype=torch.float32, device=device) y = torch.ones_like(x) z = torch.matmul(x, y.T) print(z.shape) print(z.device)
# TensorFlow 2 import tensorflow as tf x = tf.random.normal([3, 4], dtype=tf.float32) y = tf.ones_like(x) z = tf.matmul(x, tf.transpose(y)) print(z.shape) print(z.device)
# JAX import jax import jax.numpy as jnp key = jax.random.PRNGKey(42) x = jax.random.normal(key, (3, 4), dtype=jnp.float32) y = jnp.ones_like(x) z = jnp.matmul(x, y.T) print(z.shape) print(jax.default_backend())

从这段代码能看出几个关键差异。

第一,随机数种子机制不同。PyTorch 和 TensorFlow 维护全局随机状态,JAX 没有全局随机状态,必须显式传入PRNGKey。这是 JAX 函数式设计的必然结果:一个函数不能偷偷依赖外部状态,否则无法保证jit编译的确定性。

第二,device的获取方式不同。PyTorch 可以tensor.device,TensorFlow 也是tensor.device,JAX 则通过jax.default_backend()查看默认后端,如果需要把数组放到特定设备上,要用jax.device_put

第三,维度布局习惯不同。图像任务中,PyTorch 默认是NCHW,TensorFlow 默认是NHWC,JAX 生态里更常见NCHW。这个差异在迁移 CNN 代码时最容易出错,后面会专门讲。

为了方便日常查表,这里再列一组常用操作:

操作PyTorchTensorFlowJAX
创建全零数组torch.zeros((3, 4))tf.zeros((3, 4))jnp.zeros((3, 4))
创建随机数组torch.randn((3, 4))tf.random.normal((3, 4))jax.random.normal(key, (3, 4))
类型转换tensor.to(torch.float16)tf.cast(tensor, tf.float16)tensor.astype(jnp.float16)
形状变化tensor.reshape(2, 6)tf.reshape(tensor, (2, 6))tensor.reshape(2, 6)
设备迁移tensor.to('cuda')tf.device('/GPU:0')不直接绑定设备

如果你之前只写过 PyTorch,看到 JAX 的随机数会很不习惯,但这恰恰是理解 JAX 的入口。JAX 的哲学是“函数式 + 显式数据流”,任何有副作用的行为都要被收拢到边界上。

4. 自动微分 API 对比:backward、GradientTape 与 grad

自动微分是深度学习框架的核心能力。PyTorch 用动态计算图autograd,TensorFlow 2 用tf.GradientTape,JAX 用jax.grad。三者解决的问题相同,但思考模型完全不同。

PyTorch 的做法是在前向计算过程中动态搭建计算图,每个张量记录从哪里来、以及如何计算梯度。训练时调用loss.backward(),梯度会回传到所有requires_grad=True的张量上。这种方式非常直观,调试时可以像普通 Python 一样打断点。

TensorFlow 2 的GradientTape是一个上下文管理器。在前向计算时,它会把涉及到的可训练变量和中间操作“录下来”,然后手动调用tape.gradient(loss, model.trainable_variables)获取梯度。最核心的区别是:PyTorch 的backward是“计算图自动回传”,TensorFlow 是“从磁带中取出梯度”。

JAX 则完全不同。jax.grad是一个高阶函数转换工具,它接收一个函数,返回一个新函数。新函数计算的是原函数对第一个参数在给定点的梯度。JAX 要求传入的函数是纯函数,也就是不能修改全局变量、不能依赖外部可变状态,输入输出必须通过参数和返回值显式传递。

下面是最小示例,演示一元函数y = x^2 + 2x + 1x=3处的导数,理论值是 8。

# PyTorch import torch x = torch.tensor(3.0, requires_grad=True) y = x ** 2 + 2 * x + 1 y.backward() print(x.grad) # tensor(8.)
# TensorFlow 2 import tensorflow as tf x = tf.Variable(3.0) with tf.GradientTape() as tape: y = x ** 2 + 2 * x + 1 grad = tape.gradient(y, x) print(grad.numpy()) # 8.0
# JAX import jax.numpy as jnp from jax import grad def f(x): return x ** 2 + 2 * x + 1 print(grad(f)(3.0)) # 8.0

如果还要对比 Paddle、MindSpore、MXNet,可以看出两种风格:

# PaddlePaddle import paddle x = paddle.to_tensor(3.0, stop_gradient=False) y = x ** 2 + 2 * x + 1 y.backward() print(x.grad) # Tensor(shape=[], dtype=float32, place=..., value=8.)
# MindSpore 2.x import mindspore as ms from mindspore import Tensor def f(x): return x ** 2 + 2 * x + 1 grad_fn = ms.grad(f) print(grad_fn(Tensor(3.0, ms.float32))) # 8.0
# MXNet import mxnet as mx from mxnet import autograd, nd x = nd.array([3.0]) x.attach_grad() with autograd.record(): y = x ** 2 + 2 * x + 1 y.backward() print(x.grad) # [8.]

从“编程范式”的角度看,Paddle 和 MXNet 的写法更接近 PyTorch,都是显式创建带梯度属性的张量,然后backward。MindSpore 的ms.grad更接近 JAX 的函数转换风格,但也可以用GradOperation或基于nn.Cell的方式实现。

高阶求导也能体现出差异。JAX 天然支持grad(grad(f))这样的组合,因为函数转换可以任意嵌套。PyTorch 要算二阶导,需要在backward()时设置create_graph=True,并且额外维护一个计算图。这个差别在物理模拟、科学计算等需要高阶导数的场景里尤其重要。

5. 模型构建 API 对比:nn.Module、keras.Model 与 Flax Linen

模型构建是框架 API 差异最明显的地方。PyTorch 通过继承torch.nn.Module定义模型,前向方法叫forward;TensorFlow 的 Keras 接口通过继承tf.keras.Model或直接堆Sequential,前向方法叫call;JAX 核心没有模型封装,通常用 Flax 的linen.Module,前向方法直接写在__call__里。

这里有个常见的误解:很多人以为 JAX“没有深度学习框架”,其实更准确的说法是“JAX 核心只提供数值计算原语,模型层由 Flax、Haiku、Equinox 等生态库承担”。它们都构建在 JAX 的纯函数和参数数组之上,只是帮我们管理参数集合和模块初始化。

下面用一个最简单的两层 MLP 演示三巨头:

# PyTorch import torch.nn as nn class MLP(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, out_dim) ) def forward(self, x): return self.net(x)
# TensorFlow / Keras import tensorflow as tf class MLP(tf.keras.Model): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.hidden = tf.keras.layers.Dense(hidden_dim, activation='relu') self.output_layer = tf.keras.layers.Dense(out_dim) def call(self, x): return self.output_layer(self.hidden(x)) # 或者用 Sequential # model = tf.keras.Sequential([ # tf.keras.layers.Dense(hidden_dim, activation='relu'), # tf.keras.layers.Dense(out_dim) # ])
# JAX + Flax Linen 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

三个框架的差异在这里体现得很清楚:

PyTorch 用nn.Module的子类管理参数,子模块注册在self.net下,调用model.parameters()就能拿到全部可训练参数。

Keras 继承了空实现,call里用到的Dense层会被自动追踪,训练时用model.trainable_variables拿到参数。Keras 最大的优势是高层能力完整,compile + fit几乎把标准训练流程封装好了。

Flax 的写法更像是“定义计算结构”。@nn.compact装饰器允许在__call__里直接创建子层,同时自动完成参数初始化。初始化时要调用model.init(key, example_input),返回的结果是参数字典。之后前向推理用model.apply(params, x),参数和计算彻底分离。第一次看到这种代码的人会觉得“怎么这么绕”,但正是这种分离,让 JAX 可以轻松地把一个模型应用在不同设备上,也方便做模型并行。

其他框架的类名也值得记住:Paddle 用paddle.nn.Layer,重写forward;MindSpore 用mindspore.nn.Cell,重写construct;MXNet Gluon 用mxnet.gluon.nn.Block,重写forward。它们和 PyTorch 的nn.Module属于同一类设计,只是在静态图编译、自动混合精度、设备绑定等细节上有所不同。

6. 训练循环 API 对比:高层封装与手动循环

训练循环是框架之间“生产效率”差异最大的地方,也是最容易被低估的一部分。

Keras 把标准训练流程封装成了model.compile+model.fit,用户几乎不用自己写循环。PyTorch 则倾向于把控制权交给开发者,常见写法是for batch in dataloader: optimizer.zero_grad(); loss.backward(); optimizer.step()。Paddle 有paddle.Model.fit,MindSpore 有model.train,也都提供了高层训练 API,但自定义场景下还是需要理解它们的回调机制。

JAX 没有内置fit,你必须自己写一个train_step函数,并且通常用jax.jit把它编译成高效的图执行。这个过程更底层,但换来的是极致的控制和性能优化空间。

下面演示最小训练步骤。

PyTorch 手动循环:

import torch import torch.nn as nn model = MLP(28 * 28, 128, 10) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) loss_fn = nn.CrossEntropyLoss() # 假设 dataloader 产出 (x_batch, y_batch) for x_batch, y_batch in train_loader: optimizer.zero_grad() pred = model(x_batch) loss = loss_fn(pred, y_batch) loss.backward() optimizer.step()

TensorFlow / Keras 高层封装:

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

JAX + Flax + Optax 手动训练步骤:

import jax import jax.numpy as jnp import flax.linen as nn import optax model = MLP(hidden_dim=128, out_dim=10) key = jax.random.PRNGKey(0) params = model.init(key, jnp.ones((1, 28 * 28)))['params'] optimizer = optax.adam(1e-3) opt_state = optimizer.init(params) def loss_fn(params, x, y): logits = model.apply({'params': params}, x) return jnp.mean(optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=y )) @jax.jit def train_step(params, opt_state, x, y): loss, grads = jax.value_and_grad(loss_fn)(params, x, y) updates, opt_state = optimizer.update(grads, opt_state, params) params = optax.apply_updates(params, updates) return loss, params, opt_state

从这个例子可以看到 JAX 的训练循环关键点:params不是被某个 Module 持有的,而是被当作函数参数传入传出。每次train_step都返回新的参数和优化器状态,原来的参数对象不会改变。这种不可变更新模式一开始可能不习惯,但它让并行和编译变得非常安全。

如果你用 MindSpore 或 Paddle,会发现它们都在统一静态图方向做了很多工作:Paddle 的paddle.jit.to_static可以把动态图代码转成静态图,MindSpore 则从设计上强调静态编译和自动微分统一。这类框架适合在特定硬件上做高性能推理,但 API 的“哲学”和 PyTorch 不完全一致,迁移时要注意stop_gradientno_grad、梯度累积等细节。

7. 模型保存与加载 API 对比

模型保存与加载是跨框架迁移时最容易被忽视的坑。很多人以为“保存模型”就是把一个文件存下来,实际上,不同框架的保存格式包含的内容完全不同:可能是模型权重,可能是完整计算图,也可能是带优化器状态的训练检查点。

PyTorch 最常用的是torch.save(model.state_dict(), "model.pt"),加载时先创建模型实例,再load_state_dict

torch.save(model.state_dict(), "model.pt") model = MLP(28 * 28, 128, 10) state_dict = torch.load("model.pt") model.load_state_dict(state_dict)

需要注意的是,PyTorch 2.6 开始,torch.loadweights_only默认值变成了True,也就是只允许加载张量等安全对象。如果旧模型文件里用 pickle 保存过其他 Python 对象,加载时可能需要显式设置weights_only=False,但这只应该用于你完全信任的模型文件。加载陌生来源的.pt文件本身就是反序列化风险,不要随意执行。

TensorFlow / Keras 推荐保存为 SavedModel:

model.save("my_model", save_format="tf") loaded_model = tf.keras.models.load_model("my_model")

SavedModel 包含模型结构、权重和部分执行逻辑,在 TensorFlow Serving 里可以直接使用。如果只想保存权重,也可以用model.save_weights("model.weights.h5"),加载前需要先构建同样结构的模型。

JAX 没有统一的“模型文件”概念,因为模型就是参数字典。常见的做法是把paramsnp.savez保存,或者用 Flax / Orbax 的序列化工具:

import numpy as np params = model.init(key, jnp.ones((1, 28 * 28)))['params'] # 保存为 numpy 字典 np.savez("params.npz", **{k: np.asarray(v) for k, v in flatten_params(params).items()})

这个例子是示意,实际上 Flax 提供了flax.serialization.to_bytesfrom_bytes来保存参数字典。由于 JAX 模型没有“计算图”概念,保存的文件只包含参数值,恢复时需要重新定义模型结构并调用model.init或构造正确的参数字典。

Paddle 的保存方式与 PyTorch 类似:

paddle.save(model.state_dict(), "model.pdparams")

MindSpore 使用mindspore.save_checkpoint

ms.save_checkpoint(model, "model.ckpt")

MXNet Gluon 使用net.save_parameters("model.params")net.load_parameters

给一个安全提醒:加载任何模型文件之前,要确认来源可信。PyTorch 的 pickle 机制、Keras 早期 H5 文件、甚至某些框架的回调机制都曾出现过反序列化风险。不要随便从陌生网站下载“预训练模型”并直接加载。

8. 常见问题与排查:API 迁移中的高频坑

问题现象可能原因排查方式解决方案
图像维度不一致PyTorch 默认 NCHW,TensorFlow 默认 NHWC打印张量 shape,对比第一个维度使用tf.transposetorch.permute显式转换
设备不匹配报错CPU 张量与 GPU 张量参与运算检查tensor.devicetensor.device输出统一调用to(device)tf.device上下文
梯度没有更新忘记optimizer.zero_grad(),或参数没设requires_grad打印 loss 和梯度值,检查model.parameters()backward前清空梯度,确认参数可训练
JAX 的jit编译报错函数内部使用了全局可变状态查看错误信息,确认pure function要求改为显式传参,使用jax.debug.print调试
加载权重 key 不匹配PyTorch 用state_dict,Flax 用params字典打印state_dict或参数字典的 keys对 PyTorch 使用strict=False或调整键名
TensorFlowtf.function难以调试图执行阶段错误信息不直观先在 Eager 模式下跑通tf.config.run_functions_eagerly(True)临时调试
模型保存后无法加载保存的是完整模型还是权重不匹配查看文件大小和加载报错统一使用框架推荐的保存方式

这些坑里,维度顺序和自动微分模型是最影响迁移效率的两个。我的建议是:切换框架后先跑通一个最简单的全连接网络,打印每一步的输入输出 shape,不要一上来就迁移 ResNet 或 Transformer 这种大模型。等小模型验证通过,再逐步扩大范围。

9. 如何选型与迁移建议:最终判断

回到最现实的问题:到底该选哪个框架?我的判断是:

PyTorch 依然是研究社区和开源模型的主力选择。它的动态图机制对调试友好,生态完善,从最新论文到 HuggingFace 模型库都有大量 PyTorch 实现。如果你的工作重点是快速验证算法、复现论文、或者需要高度灵活的模型结构,PyTorch 是稳妥选择。

TensorFlow 在成熟企业系统里的地位仍然不可忽视。很多已经跑了好几年的生产链路、TensorFlow Serving 服务、移动端模型转换,都是围绕 TensorFlow 搭建的。如果你需要长时间维护一个稳定的推理系统,并且团队成员已经熟悉 Keras,沿用 TensorFlow 并不丢人。

JAX 适合对性能、并行和大规模训练有极致要求的场景。函数式 API 的上手曲线比 PyTorch 陡,但一旦习惯,你会发现vmappmapjit这些能力非常适合做科学计算、强化学习和大型模型训练。不过,JAX 生态的工程化组件相对分散,需要你愿意自己组装工具链。

Keras 适合作为快速建模和多后端迁移的中间层。PaddlePaddle 和 MindSpore 则要看你的部署硬件和平台需求,它们在自己的生态里确实提供了不少增强能力。MXNet 我不建议新项目继续使用,老项目维护时按照现有代码风格走即可。

如果你要做框架迁移,不要逐行翻译代码,而是先做概念映射。把“优化器、损失函数、数据加载、模型保存”这些模块独立出来,确认每个模块在两个框架中的对应关系,再动手改代码。有条件的话,先用一个固定数据集和固定随机种子建立基准,确保迁移前后指标一致,再继续扩展。

这篇文章重点对比了七个框架在张量、自动微分、模型构建、训练循环、模型保存五个维度的核心 API,并给出了迁移中的常见坑和排查思路。建议收藏备用。下一篇可以继续深入一个方向:如何用 ONNX 打通多框架部署链路,或者用 JAX 实现一个可微编程的小案例。你更想先看哪个?

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

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

立即咨询