
最近在帮团队做深度学习选型时发现一个很有意思的现象同样是“用框架训练模型”PyTorch 里创建对象叫tensorTensorFlow 里叫Tensor到了 JAX 又变成Array三个框架对“梯度”“模型”“训练循环”的抽象方式也完全不同。很多同学在 PyTorch 和 TensorFlow 之间反复横跳又听说 JAX 在科研和高性能计算圈越来越火结果不知道应该深入哪一个。本文整理了一份 7 个深度学习框架核心 API 的对比详解覆盖 PyTorch、TensorFlow、Keras 3、JAX、PaddlePaddle、MindSpore、MXNet重点拆解张量创建、自动微分、模型搭建、数据加载、图编译这几类 API 的异同并给出三个主流框架可直接运行的线性回归训练示例最后附上常见报错排查清单。无论你是刚入门深度学习还是已经用其中一个框架写过项目都可以从横向对比中获得选型参考。1. 深度学习框架核心 API 是什么为什么要对比1.1 先澄清一个概念这里的 API 不是 HTTP 接口本文讨论的 API是框架暴露给 Python 开发者的编程接口也就是你import torch、import tensorflow、import jax之后调用的一整套类、函数和调用约定。这部分和“大模型平台 API”不是一回事。近期在一些模型服务平台上经常看到api error: 400、connection lost mid-response、insufficient balance之类的报错这些属于 HTTP/JSON 服务调用层的问题和模型内部的张量运算没有直接关系。建议初学的读者先把这两个概念分开一类是“框架给我提供什么函数写模型”另一类是“我通过什么接口请求一个已经部署好的模型服务”。1.2 为什么核心 API 值得横向对比深度学习框架虽然在功能上大同小异但设计哲学差异非常大。PyTorch 走的是“命令式 动态图”路线写起来像普通 PythonTensorFlow 从静态图转向 Eager 之后又用tf.function和 Keras 把高层接口统一起来JAX 则完全采用函数式编程用grad、jit、vmap组合变换而不是维护一个“模型对象”的内部状态。如果你只学过其中一个框架换到另一个框架时最容易踩的坑不是语法不熟而是思维模式没转过来。比如 PyTorch 里你需要手动调用optimizer.zero_grad()JAX 里则因为没有可变梯度状态每次都要显式地“传入参数、返回梯度”。这些 API 差异背后是框架对自动微分和计算图执行方式的不同设计。所以本文不是要让你一次学会七个框架而是帮你建立一张“API 地图”知道每个框架在什么时候该查哪一类 API在选型和迁移时能少走弯路。2. 七个框架的定位与选型总览框架主要发起方默认执行模式核心优势典型适用场景PyTorchPyTorch 基金会Meta 主导动态图默认科研灵活、生态丰富、调试方便学术研究、快速原型、CV/NLP 项目TensorFlowGoogle动态图优先可转静态图生产部署闭环、Keras 高层 API 成熟工业落地、移动端/服务端部署Keras 3Google 等多后端高层 API同一套代码可切换 PyTorch/TF/JAX 后端快速建模、教学演示、跨后端复用JAXGoogle函数式 JIT 编译自动微分/向量化/编译组合能力强科研算法、高性能计算、强化学习PaddlePaddle百度动静统一中文文档完善、产业案例多国内产业项目、中文社区MindSpore华为动静统一昇腾硬件生态、全场景 AI 框架昇腾设备、端边云协同场景MXNetApache动态图 混合模式曾广泛用于 AWS 生态存量项目维护、历史代码阅读需要说明的是上面的“主要发起方”和框架当前的实际发展速度有关版本和社区动态变化很快尤其是 MXNet 已经进入低活跃维护阶段新项目不建议优先选择。更多精确信息请以各框架官方文档为准。3. 张量创建与基础运算 API 对比张量是深度学习框架最基本的抽象。七个框架在“创建张量”这一动作上看起来差不多但在维度约定、设备管理、梯度开关等细节上各有差别。3.1 七个框架张量创建写法# PyTorch import torch x torch.tensor([[1, 2], [3, 4]], dtypetorch.float32) print(x.shape, x.device)# TensorFlow import tensorflow as tf x tf.constant([[1, 2], [3, 4]], dtypetf.float32) print(x.shape, x.device)# JAX import jax.numpy as jnp x jnp.array([[1, 2], [3, 4]], dtypejnp.float32) print(x.shape, x.devices())# Keras 3以 TensorFlow 为后端 import keras from keras import ops x ops.convert_to_tensor([[1, 2], [3, 4]], dtypefloat32) print(ops.shape(x))# PaddlePaddle import paddle x paddle.to_tensor([[1, 2], [3, 4]], dtypefloat32) print(x.shape, x.place)# MindSpore import mindspore as ms from mindspore import Tensor from mindspore.common import dtype as mstype x Tensor([[1, 2], [3, 4]], mstype.float32) print(x.shape)# MXNet from mxnet import nd x nd.array([[1, 2], [3, 4]], dtypefloat32) print(x.shape, x.context)如果你把这七段代码并排摆放会发现一个规律torch.tensor、tf.constant、jnp.array、paddle.to_tensor、Tensor、nd.array本质上都是“把 Python 列表转成框架张量”不同点在于默认设备和 dtype 的处理策略。3.2 关键差异点操作PyTorchTensorFlowJAXPaddlePaddleMindSpore创建张量torch.tensortf.constantjnp.arraypaddle.to_tensormindspore.Tensor张量形状x.shapex.shapex.shapex.shapex.shape梯度开关requires_gradTruetf.Variablegrad按参数计算stop_gradientFalse参与grad的输入自动求导默认设备管理CPU 默认.to(cuda)GPU 自动分配为主显式jax.device_putpaddle.set_devicems.set_context(device_targetAscend)一个容易忽略的细节是 JAX 遵循“函数式纯量”原则jnp.array创建出来的数组是不可变的你不能像 PyTorch 那样直接x[0, 0] 1原地修改。所有看似“修改”的操作实际上都是生成一个新数组。这在写训练循环时需要特别注意。TensorFlow 里则有“张量不可变、变量可变”的区分tf.constant用于创建不可变张量模型参数需要放在tf.Variable里PyTorch 则没有单独区分参数和普通张量统一由requires_grad控制。这也是两个框架 API 设计上一个很典型的差异。4. 自动微分核心 API 对比自动微分是深度学习框架最核心的能力。几乎每个框架都提供“根据计算图自动求梯度”的机制但触发方式差异很大。4.1 同一个函数五种求导写法下面用同一个函数 ( f(x) x^2 2x ) 为例在五个框架中分别计算 ( x3 ) 处的导数。理论上结果都应该是 8。# PyTorch反向传播 requires_grad import torch x torch.tensor(3.0, requires_gradTrue) y x ** 2 2 * x y.backward() print(x.grad) # tensor(8.)# TensorFlowGradientTape 上下文 import tensorflow as tf x tf.Variable(3.0) with tf.GradientTape() as tape: y x ** 2 2 * x grad tape.gradient(y, x) print(grad.numpy()) # 8.0# JAXgrad 函数变换 import jax import jax.numpy as jnp def f(x): return x ** 2 2 * x print(jax.grad(f)(3.0)) # 8.0# PaddlePaddle与 PyTorch 类似的动态图 import paddle x paddle.to_tensor(3.0, stop_gradientFalse) y x ** 2 2 * x y.backward() print(x.grad) # Tensor(shape[], dtypefloat32, placeCPUPlace, value8)# MindSporegrad 函数式接口 import mindspore as ms def f(x): return x ** 2 2 * x grad_fn ms.grad(f) print(grad_fn(ms.Tensor(3.0, ms.float32))) # 8.04.2 三种自动微分风格的本质区别从上面的代码可以清晰看到三类风格PyTorch 和 PaddlePaddle 是“命令式反向传播”。张量内部会记录计算图你正常写运算最后调用backward()梯度从输出端反向传播到叶子张量。这种方式直观适合调试。TensorFlow 用GradientTape上下文管理器显式“录制”一段计算过程然后在上下文中调用tape.gradient提取梯度。这比 PyTorch 多一步显式管理但也让“哪些计算参与求导”变得边界清晰。JAX 和 MindSpore 的函数式grad则完全不同。jax.grad(f)返回一个新函数这个新函数对传入参数求导原函数内部不应该有全局可变状态。如果你在grad内部尝试修改 Python 列表或使用 NumPy 数组运算很容易触发TracerArrayConversionError因为 JAX 需要把计算过程转化为可追踪的轨迹。5. 模型定义与网络搭建 API 对比模型定义是框架 API 差异最明显的地方。PyTorch 用nn.Module子类TensorFlow 用tf.keras.Model子类JAX 生态里常用 Flax 的nn.ModulePaddlePaddle 用nn.LayerMindSpore 用nn.Cell。5.1 同一个两层 MLP不同框架的写法# PyTorch import torch.nn as nn class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(4, 16) self.relu nn.ReLU() self.fc2 nn.Linear(16, 1) def forward(self, x): return self.fc2(self.relu(self.fc1(x)))# TensorFlow import tensorflow as tf class MLP(tf.keras.Model): def __init__(self): super().__init__() self.fc1 tf.keras.layers.Dense(16, activationrelu) self.fc2 tf.keras.layers.Dense(1) def call(self, x): return self.fc2(self.fc1(x))# Keras 3Sequential 高层接口 import keras model keras.Sequential([ keras.layers.Dense(16, activationrelu), keras.layers.Dense(1), ])# JAX Flax import flax.linen as nn class MLP(nn.Module): nn.compact def __call__(self, x): x nn.Dense(16)(x) x nn.relu(x) x nn.Dense(1)(x) return x# PaddlePaddle import paddle.nn as nn class MLP(nn.Layer): def __init__(self): super().__init__() self.fc1 nn.Linear(4, 16) self.fc2 nn.Linear(16, 1) def forward(self, x): return self.fc2(nn.functional.relu(self.fc1(x)))# MindSpore import mindspore.nn as nn class MLP(nn.Cell): def __init__(self): super().__init__() self.fc1 nn.Dense(4, 16) self.fc2 nn.Dense(16, 1) self.relu nn.ReLU() def construct(self, x): return self.fc2(self.relu(self.fc1(x)))# MXNet Gluon from mxnet.gluon import nn net nn.Sequential() net.add(nn.Dense(16, activationrelu)) net.add(nn.Dense(1))5.2 名称不同思路相通观察上面代码你会发现除了 JAX/Flax 的__call__和 MindSpore 的construct命名特殊其他框架几乎都是“构造函数里定义层前向方法里连接层”只是方法名从forward变成call。JAX Flax 的模型定义是最“函数式”的nn.Dense(16)并不绑定一个具体的参数数组参数是在初始化阶段由 PRNGKey 生成的。这也意味着 JAX 风格的模型不能像 PyTorch 那样“创建完就直接用”你必须先调用init拿到参数再通过apply执行前向计算。这是新手从 PyTorch 转向 JAX 时最不习惯的点。6. 数据加载与训练循环 API 对比数据加载也是框架 API 的差异化重灾区。PyTorch 的DataLoader、TensorFlow 的tf.data.Dataset、JAX 的纯 NumPy 批处理、PaddlePaddle 的paddle.io.DataLoader、MindSpore 的mindspore.dataset都承担着“把原始数据组织成 batch”的职责。6.1 PyTorch DataLoader 示例import torch from torch.utils.data import DataLoader, TensorDataset x torch.randn(1000, 4) y torch.randn(1000, 1) dataset TensorDataset(x, y) loader DataLoader(dataset, batch_size32, shuffleTrue) for batch_x, batch_y in loader: print(batch_x.shape, batch_y.shape) break6.2 TensorFlow tf.data 示例import tensorflow as tf x tf.random.normal((1000, 4)) y tf.random.normal((1000, 1)) dataset tf.data.Dataset.from_tensor_slices((x, y)) dataset dataset.shuffle(1000).batch(32) for batch_x, batch_y in dataset.take(1): print(batch_x.shape, batch_y.shape)6.3 JAX 通常直接配合 NumPy 做批处理import jax.numpy as jnp key jax.random.PRNGKey(0) x jax.random.normal(key, (1000, 4)) y jax.random.normal(key, (1000, 1)) batch_size 32 num_batches x.shape[0] // batch_size for i in range(num_batches): batch_x x[i * batch_size:(i 1) * batch_size] batch_y y[i * batch_size:(i 1) * batch_size] print(batch_x.shape, batch_y.shape)从 API 设计上看PyTorch 的DataLoader最“重”因为它内置了多进程加载、shuffle、采样器等机制TensorFlow 的tf.data是“惰性流水线”强调高性能输入管道JAX 没有独立的 DataLoader生态里通常用tf.data或PyTorch DataLoader先把数据准备好再转成 JAX 数组或者直接基于 NumPy 切片。训练循环的差异也一样PyTorch 需要手动写optimizer.zero_grad()、loss.backward()、optimizer.step()三步TensorFlow Keras 可以用model.fit()一行完成也可以用GradientTape自定义循环JAX 则需要完全手动实现“前向计算、求梯度、更新参数”的过程但配合jax.jit可以获得不错的性能。7. 动态图、静态图与编译 API 对比深度学习框架的执行模式直接决定代码的组织方式。PyTorch 默认动态图TensorFlow 加入了tf.functionJAX 默认通过jit编译PaddlePaddle 和 MindSpore 则提供“动静统一”的方案。7.1 三种编译 API 的最小示例# PyTorch 2.xtorch.compile 编译模型 import torch model torch.nn.Linear(4, 1) model torch.compile(model) x torch.randn(8, 4) y model(x) print(y.shape)# TensorFlowtf.function 将 Python 函数编译为计算图 import tensorflow as tf tf.function def predict(x): return tf.linalg.matmul(x, w) b w tf.Variable(tf.random.normal((4, 1))) b tf.Variable(tf.zeros((1,))) x tf.random.normal((8, 4)) print(predict(x).shape)# JAXjax.jit 编译计算函数 import jax import jax.numpy as jnp def predict(w, b, x): return x w b w jnp.zeros((4, 1)) b jnp.zeros((1,)) x jnp.ones((8, 4)) fast_predict jax.jit(predict) print(fast_predict(w, b, x).shape)7.2 为什么需要编译动态图模式下每一行 Python 代码都会触发一次张量运算灵活但开销大静态图模式先把整个计算过程描述成一张图再交给底层编译器优化并执行在重复训练和推理时性能更好。JAX 的设计思路是最彻底的函数默认不编译但你可以用jit把任意纯函数编译成高效的 XLA 计算。也正因如此JAX 里有一句名言“如果你想性能好就把计算写进jit里”。PyTorch 的torch.compile是近几年重点演进的方向它试图在不改变动态图编程体验的前提下通过编译加速模型。实际使用中torch.compile对 GPU 显存的占用、编译时间都有影响并不是所有模型都能无脑加速需要针对性验证。8. 完整实战Python 环境准备与三框架线性回归对比下面用一个真实可运行的线性回归任务完整对比 PyTorch、TensorFlow、JAX 三个框架的训练循环。数据统一为 ( y 2x 1 )添加少量高斯噪声。8.1 环境准备建议建议使用 conda 创建独立虚拟环境避免框架依赖互相污染。安装命令以官方文档为准不同 CUDA 版本对应不同的安装命令。例如 PyTorch 2.x 在安装时要注意选择与本地 CUDA 驱动匹配的版本TensorFlow 2.18 安装时也要确认 Python 版本和 GPU 支持情况。先跑 CPU 版本再补 GPU 版本是更稳妥的路径。conda create -n dl-compare python3.10 -y conda activate dl-compare # PyTorch、TensorFlow、JAX 分开安装优先参考官方安装页 pip install torch pip install tensorflow pip install jax jaxlib注意这不是一条命令同时装完的环境实际项目中建议按需安装本文为了演示三个框架可以共存。8.2 PyTorch 训练示例import torch import torch.nn as nn torch.manual_seed(0) x torch.linspace(-1, 1, 100).reshape(-1, 1) y 2 * x 1 0.1 * torch.randn_like(x) model nn.Linear(1, 1) opt torch.optim.SGD(model.parameters(), lr0.1) loss_fn nn.MSELoss() for epoch in range(200): opt.zero_grad() loss loss_fn(model(x), y) loss.backward() opt.step() print(fw{model.weight.item():.3f}, b{model.bias.item():.3f})输出接近w2.000, b1.0008.3 TensorFlow 训练示例import tensorflow as tf tf.random.set_seed(0) x tf.linspace(-1.0, 1.0, 100)[:, None] y 2 * x 1 0.1 * tf.random.normal(x.shape) model tf.keras.Sequential([tf.keras.layers.Dense(1)]) opt tf.keras.optimizers.SGD(0.1) loss_fn tf.keras.losses.MeanSquaredError() for epoch in range(200): with tf.GradientTape() as tape: loss loss_fn(y, model(x)) grads tape.gradient(loss, model.trainable_variables) opt.apply_gradients(zip(grads, model.trainable_variables)) kernel model.layers[0].kernel.numpy()[0][0] bias model.layers[0].bias.numpy()[0] print(fw{kernel:.3f}, b{bias:.3f})输出接近w2.000, b1.0008.4 JAX 训练示例import jax import jax.numpy as jnp key jax.random.PRNGKey(0) x jnp.linspace(-1.0, 1.0, 100).reshape(-1, 1) y 2 * x 1 0.1 * jax.random.normal(key, x.shape) def predict(w, b, x): return x w b def loss_fn(w, b, x, y): return jnp.mean((predict(w, b, x) - y) ** 2) w jnp.zeros((1, 1)) b jnp.zeros((1,)) lr 0.1 for _ in range(200): grad_w, grad_b jax.grad(loss_fn, argnums(0, 1))(w, b, x, y) w w - lr * grad_w b b - lr * grad_b print(float(w[0][0]), float(b[0]))输出接近2.000 1.0008.5 三个训练循环的差异总结对比项PyTorchTensorFlowJAX参数对象nn.Linear内部参数tf.Variable普通数组jnp.array求梯度方式loss.backward()tape.gradient(loss, vars)jax.grad(loss_fn, argnums...)参数更新optimizer.step()optimizer.apply_gradients(...)手动w w - lr * grad_w随机种子torch.manual_seedtf.random.set_seedjax.random.PRNGKey从表中可以看出JAX 的更新是最显式的也最接近数学公式PyTorch 和 TensorFlow 把更新细节封装在优化器里但换来了更少的样板代码。9. 常见问题与排查思路在实际使用中框架安装和运行报错是最消耗时间的环节。下面整理几个高频问题。问题现象常见原因解决思路torch.load报UnpicklingError提示weights_only参数PyTorch 2.6 开始torch.load默认weights_onlyTrue旧模型文件无法直接加载对可信的模型文件使用torch.load(path, weights_onlyFalse)优先用官方推荐的保存格式TensorFlow 找不到 GPUCUDA、cuDNN、TensorFlow 三者版本不匹配用tf.config.list_physical_devices(GPU)检查根据官方安装页对照版本tf.function每次重新编译训练变慢输入 shape 或 dtype 每次变化触发 retracing固定输入维度或对可变维度使用 padding 和 maskJAX 报TracerArrayConversionError在jit/grad内部混用了 NumPy 数组或 Python 控制流全部改用jnp操作条件分支改用jax.lax.condJAX 反复编译耗时传入jit的数组 shape 或 dtype 不稳定确保同一函数内张量维度稳定避免每次生成不同 shapePaddlePaddle 报 CUDNN 相关错误安装版本与 CUDA 版本不一致用对应 CUDA 版本的安装命令重新安装确认nvidia-smi驱动版本MindSpore 图模式报错、动态模式正常图编译对 Python 语法支持有限先用 PYNATIVE 模式调试再切到 GRAPH 模式验证MXNet 在 Windows 上导入失败老版本与新版 Python 不兼容使用官方 Docker 镜像或降级到匹配的 Python 版本9.1 排查清单遇到框架 API 或环境报错按下面顺序排查先确认 Python 版本、CUDA 版本、框架版本三者是否匹配。用一段最小代码复现逐步缩小问题范围。如果是梯度相关报错检查是否在with或backward()作用域之外访问了梯度。如果是 JAX检查是否在jit内部使用了不纯的函数或 NumPy 操作。如果是模型加载报错优先怀疑序列化格式或版本变化而不是模型本身。10. 最佳实践与工程建议10.1 环境管理强烈建议每个项目使用独立虚拟环境并把依赖版本写入requirements.txt或pyproject.toml。深度学习框架之间依赖的 NumPy、protobuf、CUDA 版本经常冲突集中安装容易互相污染。用一个简单的环境隔离能省掉大量排查时间。10.2 可复现性训练实验要固定随机种子。PyTorch 用torch.manual_seed(0)TensorFlow 用tf.random.set_seed(0)JAX 要显式创建jax.random.PRNGKey(0)并在每次采样时 split。JAX 的随机数设计最严格因为它要求函数无状态但这反而让复现变得更加确定。10.3 模型保存与加载安全模型权重本质上是序列化数据可能包含恶意构造的对象。在 PyTorch 2.6 中官方将torch.load的weights_only默认改为True正是出于安全考虑。实际工程中只从可信来源加载模型权重不要加载来路不明的.pth或.pt文件。能使用官方推荐的torch.save(model.state_dict(), ...)就尽量不用整个模型对象。TensorFlow 推荐使用SavedModel格式导出方便部署和服务化。JAX 生态优先考虑使用safetensors或orbax保存参数数组。10.4 训练流程规范不要一上来就跑全量数据。建议先在小规模子集上验证 Loss 能下降再切到完整数据训练过程中要定期保存 checkpoint记录每个 epoch 的 loss、学习率、GPU 显存等指标。混合精度训练可以显著提升吞吐但要注意损失缩放和数值稳定性。10.5 部署思维选框架如果你是纯科研场景PyTorch 的灵活性和调试体验是最好的如果你是做模型服务TensorFlow 的 SavedModel、TensorFlow Serving 链路更成熟如果你想深入研究自动微分和编译器JAX 的函数式设计会让你看到更多底层细节。PaddlePaddle 和 MindSpore 在中文资料、国内云环境和特定硬件上的优势也值得在具体业务中重点关注。11. 下一步学习建议如果只记一句话选型建议是这样的科研原型和快速迭代用 PyTorch追求极致性能、愿意研究函数式编程和编译器细节选 JAX需要把模型推上生产服务、且团队熟悉 Keras 高层接口选 TensorFlow。PaddlePaddle 和 MindSpore 则在中文文档、产业项目和特定硬件生态上各有优势遇到真实需求时可以重点考察。接下来建议动手做三件事第一在这个对比案例基础上把线性回归换成 MNIST 手写数字分类分别用 PyTorch、TensorFlow、JAX 各写一遍感受差异。第二把同一个模型分别用torch.compile、tf.function、jax.jit编译比较训练速度和显存占用。第三去读官方文档里“自动微分”和“序列化”两个章节这两个模块决定了你对框架 API 的理解深度。如果本文对你有帮助可以收藏备用。动手写代码时遇到具体报错欢迎在评论区带上你的版本号、CUDA 版本和完整报错日志一起讨论。