
如果你在小团队里既要做算法实验又要接线上推理大概率会被几个框架的 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 用Layer和Model跟踪参数JAX 核心库没有“模型”概念参数必须显式放在函数签名里常用Flax、Haiku、Equinox这类生态库来辅助。这个差异会直接影响你的训练循环怎么写。第三生态和部署链路不同。PyTorch 在科研社区最活跃TensorFlow 在成熟的企业系统里存量巨大JAX 在大模型训练和高性能科学计算上增长很快。选框架从来不是“谁更好”的问题而是“谁更适合你当前的场景”的问题。所以这篇文章不打算评价哪个框架“最强”而是想把它们的核心 API 放到同一张表里帮你看清楚映射关系。对技术读者来说掌握跨框架抽象能力比死记单一框架更能应对项目变化。2. 七个框架的定位与生态一张表看懂在进入代码之前先建立整体认知。七个框架分别是 PyTorch、TensorFlow、JAX、Keras、PaddlePaddle、MindSpore、MXNet。框架维护方核心编程范式主要接口入口典型场景PyTorchLinux Foundation / Meta命令式动态图torch.nn、torch.optim、torch.autograd科研实验、快速原型、TorchServe 部署TensorFlowGoogle声明式图 动态执行tf.keras、tf.data、tf.function生产环境、移动端、成熟 MLOps 链路JAXGoogle Research函数式 XLA 编译jax.numpy、jax.grad、jax.jit、jax.vmap高性能科学计算、大规模并行训练KerasGoogle 社区高层声明式keras.Model、keras.layers快速建模、多后端迁移PaddlePaddle百度动静态统一paddle.nn.Layer、paddle.to_static本地生态、全流程工业平台MindSpore华为动静态统一 / 图编译mindspore.nn.Cell、mindspore.trainAI 与科学计算融合、特定硬件平台MXNetApache动态 混合编程mxnet.nd、mxnet.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.TensorTensorFlow 叫tf.TensorJAX 叫jax.ArrayPaddle 叫paddle.TensorMindSpore 叫mindspore.TensorMXNet 叫mxnet.ndarray.NDArray。虽然名字不同但它们都要处理三个核心问题数据类型 dtype、形状 shape、设备 device。新手最容易忽略的是设备语义。PyTorch 里你能直接看到tensor.device是 CPU 还是 CUDATensorFlow 也有类似概念JAX 则默认把数组看作“不关心设备”的抽象值统一由jax.jit和jax.device_put来管理。这意味着 JAX 代码看起来更干净但如果你不熟悉它的异步调度反而可能发现显存占用异常。下面用同一段“创建 3×4 随机矩阵并计算矩阵乘法”演示三个框架的写法。# PyTorch import torch device cuda if torch.cuda.is_available() else cpu x torch.randn(3, 4, dtypetorch.float32, devicedevice) 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], dtypetf.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), dtypejnp.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.deviceTensorFlow 也是tensor.deviceJAX 则通过jax.default_backend()查看默认后端如果需要把数组放到特定设备上要用jax.device_put。第三维度布局习惯不同。图像任务中PyTorch 默认是NCHWTensorFlow 默认是NHWCJAX 生态里更常见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 用动态计算图autogradTensorFlow 2 用tf.GradientTapeJAX 用jax.grad。三者解决的问题相同但思考模型完全不同。PyTorch 的做法是在前向计算过程中动态搭建计算图每个张量记录从哪里来、以及如何计算梯度。训练时调用loss.backward()梯度会回传到所有requires_gradTrue的张量上。这种方式非常直观调试时可以像普通 Python 一样打断点。TensorFlow 2 的GradientTape是一个上下文管理器。在前向计算时它会把涉及到的可训练变量和中间操作“录下来”然后手动调用tape.gradient(loss, model.trainable_variables)获取梯度。最核心的区别是PyTorch 的backward是“计算图自动回传”TensorFlow 是“从磁带中取出梯度”。JAX 则完全不同。jax.grad是一个高阶函数转换工具它接收一个函数返回一个新函数。新函数计算的是原函数对第一个参数在给定点的梯度。JAX 要求传入的函数是纯函数也就是不能修改全局变量、不能依赖外部可变状态输入输出必须通过参数和返回值显式传递。下面是最小示例演示一元函数y x^2 2x 1在x3处的导数理论值是 8。# PyTorch import torch x torch.tensor(3.0, requires_gradTrue) 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_gradientFalse) y x ** 2 2 * x 1 y.backward() print(x.grad) # Tensor(shape[], dtypefloat32, place..., value8.)# 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_graphTrue并且额外维护一个计算图。这个差别在物理模拟、科学计算等需要高阶导数的场景里尤其重要。5. 模型构建 API 对比nn.Module、keras.Model 与 Flax Linen模型构建是框架 API 差异最明显的地方。PyTorch 通过继承torch.nn.Module定义模型前向方法叫forwardTensorFlow 的 Keras 接口通过继承tf.keras.Model或直接堆Sequential前向方法叫callJAX 核心没有模型封装通常用 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, activationrelu) 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, activationrelu), # 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重写forwardMindSpore 用mindspore.nn.Cell重写constructMXNet Gluon 用mxnet.gluon.nn.Block重写forward。它们和 PyTorch 的nn.Module属于同一类设计只是在静态图编译、自动混合精度、设备绑定等细节上有所不同。6. 训练循环 API 对比高层封装与手动循环训练循环是框架之间“生产效率”差异最大的地方也是最容易被低估的一部分。Keras 把标准训练流程封装成了model.compilemodel.fit用户几乎不用自己写循环。PyTorch 则倾向于把控制权交给开发者常见写法是for batch in dataloader: optimizer.zero_grad(); loss.backward(); optimizer.step()。Paddle 有paddle.Model.fitMindSpore 有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(), lr1e-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( optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy] ) model.fit(train_dataset, epochs5)JAX Flax Optax 手动训练步骤import jax import jax.numpy as jnp import flax.linen as nn import optax model MLP(hidden_dim128, out_dim10) 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( logitslogits, labelsy )) 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_gradient、no_grad、梯度累积等细节。7. 模型保存与加载 API 对比模型保存与加载是跨框架迁移时最容易被忽视的坑。很多人以为“保存模型”就是把一个文件存下来实际上不同框架的保存格式包含的内容完全不同可能是模型权重可能是完整计算图也可能是带优化器状态的训练检查点。PyTorch 最常用的是torch.save(model.state_dict(), model.pt)加载时先创建模型实例再load_state_dicttorch.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.load的weights_only默认值变成了True也就是只允许加载张量等安全对象。如果旧模型文件里用 pickle 保存过其他 Python 对象加载时可能需要显式设置weights_onlyFalse但这只应该用于你完全信任的模型文件。加载陌生来源的.pt文件本身就是反序列化风险不要随意执行。TensorFlow / Keras 推荐保存为 SavedModelmodel.save(my_model, save_formattf) loaded_model tf.keras.models.load_model(my_model)SavedModel 包含模型结构、权重和部分执行逻辑在 TensorFlow Serving 里可以直接使用。如果只想保存权重也可以用model.save_weights(model.weights.h5)加载前需要先构建同样结构的模型。JAX 没有统一的“模型文件”概念因为模型就是参数字典。常见的做法是把params用np.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_bytes和from_bytes来保存参数字典。由于 JAX 模型没有“计算图”概念保存的文件只包含参数值恢复时需要重新定义模型结构并调用model.init或构造正确的参数字典。Paddle 的保存方式与 PyTorch 类似paddle.save(model.state_dict(), model.pdparams)MindSpore 使用mindspore.save_checkpointms.save_checkpoint(model, model.ckpt)MXNet Gluon 使用net.save_parameters(model.params)和net.load_parameters。给一个安全提醒加载任何模型文件之前要确认来源可信。PyTorch 的 pickle 机制、Keras 早期 H5 文件、甚至某些框架的回调机制都曾出现过反序列化风险。不要随便从陌生网站下载“预训练模型”并直接加载。8. 常见问题与排查API 迁移中的高频坑问题现象可能原因排查方式解决方案图像维度不一致PyTorch 默认 NCHWTensorFlow 默认 NHWC打印张量 shape对比第一个维度使用tf.transpose或torch.permute显式转换设备不匹配报错CPU 张量与 GPU 张量参与运算检查tensor.device或tensor.device输出统一调用to(device)或tf.device上下文梯度没有更新忘记optimizer.zero_grad()或参数没设requires_grad打印 loss 和梯度值检查model.parameters()在backward前清空梯度确认参数可训练JAX 的jit编译报错函数内部使用了全局可变状态查看错误信息确认pure function要求改为显式传参使用jax.debug.print调试加载权重 key 不匹配PyTorch 用state_dictFlax 用params字典打印state_dict或参数字典的 keys对 PyTorch 使用strictFalse或调整键名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 陡但一旦习惯你会发现vmap、pmap、jit这些能力非常适合做科学计算、强化学习和大型模型训练。不过JAX 生态的工程化组件相对分散需要你愿意自己组装工具链。Keras 适合作为快速建模和多后端迁移的中间层。PaddlePaddle 和 MindSpore 则要看你的部署硬件和平台需求它们在自己的生态里确实提供了不少增强能力。MXNet 我不建议新项目继续使用老项目维护时按照现有代码风格走即可。如果你要做框架迁移不要逐行翻译代码而是先做概念映射。把“优化器、损失函数、数据加载、模型保存”这些模块独立出来确认每个模块在两个框架中的对应关系再动手改代码。有条件的话先用一个固定数据集和固定随机种子建立基准确保迁移前后指标一致再继续扩展。这篇文章重点对比了七个框架在张量、自动微分、模型构建、训练循环、模型保存五个维度的核心 API并给出了迁移中的常见坑和排查思路。建议收藏备用。下一篇可以继续深入一个方向如何用 ONNX 打通多框架部署链路或者用 JAX 实现一个可微编程的小案例。你更想先看哪个