尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

PyTorch、TensorFlow、JAX三大框架API对比:7个核心维度讲透迁移与选型

PyTorch、TensorFlow、JAX三大框架API对比:7个核心维度讲透迁移与选型 之前在业务迭代里最常遇到的一件事就是同一个模型用 PyTorch 调通了切到同事的 TensorFlow 工程又要重写一遍论文开源代码用了 JAX想去复现又发现自动微分、优化器、数据加载的写法全都不一样。三个框架都说自己是“张量 自动微分 优化器”但真正落到代码上几乎没有任何一行能直接复用。网上教程大多只讲单个框架很少把开发者真正需要的 API 差异放到同一张工作台上去比较。这篇文章就用7 个核心 API 维度把 PyTorch、TensorFlow、JAX 放到一起做横向对比从张量创建、自动微分、网络层定义到数据加载、训练循环、模型保存与部署。文章最后还会用一个完整的二分类任务分别跑通三个框架的代码。无论是刚入门深度学习、正在做框架选型还是需要把项目从一套框架迁移到另一套框架这篇文章都能给你一份比较完整的参考。1. 背景与核心概念1.1 三大框架的定位差异先说结论PyTorch、TensorFlow、JAX 不是三个“同款不同名”的库它们在设计哲学上有本质区别。PyTorch 的核心设计是命令式动态图。你写一行代码它立即执行可以随时打印中间张量、打断调试很符合 Python 开发者的直觉。PyTorch 在学术研究、论文复现、快速原型阶段使用率非常高而且生态很丰富Hugging Face Transformers 等主流模型库都优先支持 PyTorch。TensorFlow 的定位则更偏“完整生产系统”。它从最初的计算图模式逐步演进现在主推 Keras 高层 API同时保留了tf.function、SavedModel、TensorFlow Serving、TensorFlow Lite 等一系列部署链路。在企业已有基础设施、服务化部署、移动端场景下TensorFlow 的存量项目和配套工具仍然非常多。JAX 是三者中最特别的一个。它不直接提供类似 Keras 的全套训练封装而是提供jax.numpy、jax.grad、jax.jit、jax.vmap等基础变换能力配合 Flax、Optax 等库完成模型定义和优化。JAX 的“函数式变换”风格让它在大规模并行、科学计算、强化学习、以及需要自定义编译优化的研究项目中非常受欢迎。1.2 为什么从 API 维度对比很多框架选型文章喜欢比性能、比社区热度、比部署工具但开发者上手时最先接触到的其实是 API。API 设计决定了你写代码的思维方式PyTorch 让你“面向对象 命令式”地写网络和训练循环TensorFlow 让你“先定义层再compile fit”地训练JAX 让你“把计算写成纯函数再交给 grad/jit 去变换”。如果在概念层面理解了这三个框架各自的 API 习惯再从 PyTorch 切到 JAX或者从 TensorFlow 迁到 PyTorch就不会那么痛苦。1.3 7 个核心维度总览下面 7 个维度是平时写模型几乎绕不开的 API 入口也是框架迁移时改动最大的部分编号API 维度一句话说明1张量创建与基础运算数据的基本载体和运算方式2自动微分梯度计算的核心机制3网络层定义模型结构如何组织4损失函数与优化器监督学习的两大基础组件5数据加载与预处理数据如何进入训练循环6训练循环每个 batch 如何更新参数7模型保存、加载与部署训练产物如何持久化和上线2. 环境准备与版本说明2.1 推荐环境本文示例以常见环境和版本为例重点演示 API 的写法差异具体版本需要根据你的实际环境调整。推荐使用 Anaconda 或 Miniconda 创建独立虚拟环境避免多个项目之间的 Python 和 CUDA 依赖互相污染。conda create -n dl-compare python3.10 -y conda activate dl-compare pip install --upgrade pipPython 版本建议使用 3.10 或 3.11。如果你的机器有 NVIDIA GPU还需要提前装好显卡驱动并确认驱动支持的 CUDA 版本建议用nvidia-smi查看。2.2 三个框架的安装命令PyTorch 安装时要注意选择与本地 CUDA 匹配的版本。这里给出 CPU 版的示例GPU 版建议到 PyTorch 官网生成对应命令pip install torch torchvision torchaudio如果你搜索过“pytorch 安装教程 gpu”“pytorch 官网”之类的内容会发现 GPU 版的核心其实是 CUDA 版本要对上。以 CUDA 12.1 为例PyTorch 官网会给出类似下面的命令实际以官网生成结果为准pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121TensorFlow 的安装相对简单CPU 版直接安装即可pip install tensorflow如果你关心版本号目前“tensorflow 2.18 安装”已经是比较常见的搜索词。需要注意 TensorFlow 对 CUDA 和 cuDNN 有严格的配套要求GPU 环境里如果版本不匹配往往会出现运行时无法加载动态库的错误。Windows 下建议优先考虑 WSL2 或者 Linux 环境。JAX 的安装分为 CPU 和 GPU 两种# CPU 版 pip install jax jaxlib # GPU 版以 CUDA 12 为例 pip install jax[cuda12]JAX 的神经网络层一般配合 Flax 使用优化器通常使用 Optax所以还要安装pip install flax optax这里有朋友会问JAX 本身不是不提供高层 API 吗没错JAX 主张“核心库保持精简”通过 Flax、Optax、TensorFlow Datasets 等生态库补齐模型定义、优化器和数据管道。这是 JAX 和 PyTorch、TensorFlow 一个非常大的设计差异。2.3 版本与 CUDA 注意事项三个框架的版本变化都比较快安装时有几个共通的建议先确认nvidia-smi显示的驱动最高 CUDA 版本再决定安装哪个 CUDA 配套包PyTorch、TensorFlow、JAX 对 CUDA 的要求并不是越新越好而是要与 cuDNN、驱动都兼容不要为了“必须装最新版”而选择刚发布的版本生产项目建议锁定小版本比如torch2.5.1这类明确版本号如果三套框架装在同一个环境依赖冲突概率会变大建议分开创建虚拟环境。3. 7 个核心 API 维度横向对比3.1 张量创建与基础运算张量是深度学习中数据的基本载体。三个框架都提供了对应 API但细节上有不少差异。PyTorch 创建张量非常直接import torch a torch.tensor([1.0, 2.0, 3.0]) # 从 Python 列表创建 b torch.zeros(2, 3) # 全 0 张量 c torch.randn(2, 3) # 标准正态随机张量 print(a 1)PyTorch 的 Tensor 是可变对象你可以直接在原有张量上做inplace操作比如x.add_(1)。TensorFlow 默认是 Eager 模式动态图和 PyTorch 的书写习惯比较接近import tensorflow as tf a tf.constant([1.0, 2.0, 3.0]) b tf.zeros((2, 3)) c tf.random.normal((2, 3)) print(a 1)不过 TensorFlow 的张量默认不可变没有类似 PyTorchadd_的 inplace 写法。JAX 使用jax.numpy来创建张量接口和 NumPy 很相似import jax.numpy as jnp a jnp.array([1.0, 2.0, 3.0]) b jnp.zeros((2, 3)) c jnp.ones((2, 3)) print(a 1)JAX 数组是完全不可变的。任何看起来像“原地修改”的操作实际上都是生成新数组。这个特性对新手来说需要适应但它和 JAX 的jit编译模型是配套的只有数据不可变才能安全地做函数变换和并行计算。3.2 自动微分机制自动微分是最影响框架使用体验的部分。PyTorch 的思路是“在反向传播路径上记录梯度”你在任意张量上设置requires_gradTrue经过前向计算后调用backward()梯度就会自动累积到张量的grad属性上。import torch x torch.tensor(2.0, requires_gradTrue) y x ** 2 3 * x y.backward() print(x.grad) # 2*x 3 7TensorFlow 使用tf.GradientTape做梯度记录代码结构有点像“上下文管理器”import tensorflow as tf x tf.Variable(2.0) with tf.GradientTape() as tape: y x ** 2 3 * x grad tape.gradient(y, x) print(grad.numpy()) # 7.0JAX 的设计完全不同。它不维护任何“反向传播的内部状态”而是把梯度计算看作一个函数变换你写一个纯函数fn(params, x)然后用jax.grad或jax.value_and_grad把它变成“能计算梯度的函数”。import jax import jax.numpy as jnp def fn(x): return x ** 2 3 * x grad_fn jax.grad(fn) print(grad_fn(2.0)) # 7.0PyTorch 和 TensorFlow 的自动微分是“先执行图再反向传播”JAX 则是“定义纯函数再用变换生成梯度函数”。这种函数式思维是 JAX 和另外两者最核心的区别。3.3 网络层定义PyTorch 使用nn.Module定义网络。它的风格是“类 子模块 forward方法”import torch.nn as nn class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(2, 16) self.fc2 nn.Linear(16, 1) def forward(self, x): x torch.relu(self.fc1(x)) return self.fc2(x)TensorFlow 最常用的是 Keras 的Sequential和函数式 API。简单模型可以直接堆层import tensorflow as tf model tf.keras.Sequential([ tf.keras.layers.Dense(16, activationrelu, input_shape(2,)), tf.keras.layers.Dense(1) ])如果需要复杂结构Keras 还提供Model子类化方式继承tf.keras.Model并实现call。JAX 本身没有网络层概念一般用 Flax 来定义。Flax 的写法看起来很像 PyTorch但底层逻辑是纯函数式的import flax.linen as nn class MLP(nn.Module): nn.compact def __call__(self, x): x nn.Dense(features16)(x) x nn.relu(x) x nn.Dense(features1)(x) return x注意 Flax 模型没有__init__里保存参数参数是在首次调用model.init(prng_key, x)时根据输入 shape 自动生成的。3.4 损失函数与优化器PyTorch 把损失函数放在torch.nn或torch.nn.functional中优化器放在torch.optim中import torch.optim as optim loss_fn nn.BCEWithLogitsLoss() optimizer optim.Adam(model.parameters(), lr0.01)TensorFlow 的对应 API 在tf.keras.losses和tf.keras.optimizers中loss_fn tf.keras.losses.BinaryCrossentropy(from_logitsTrue) optimizer tf.keras.optimizers.Adam(learning_rate0.01)这里from_logitsTrue表示损失函数内部自动加 Sigmoid避免数值不稳定对应 PyTorch 的BCEWithLogitsLoss。JAX 不直接提供“损失函数对象”而是推荐你写普通 Python 函数再用 Optax 管理优化器import optax def loss_fn(params, x, y): logits model.apply(params, x) return optax.sigmoid_binary_cross_entropy(logits, y).mean() optimizer optax.adam(learning_rate0.01) opt_state optimizer.init(params)JAX 的损失函数就是普通函数参数params是显式传入的。这个“所有状态都摆在明面上”的风格是函数式框架的特征。3.5 数据加载与预处理PyTorch 的数据管道是Dataset DataLoader。你可以用现成TensorDataset也可以自定义Datasetfrom torch.utils.data import TensorDataset, DataLoader dataset TensorDataset(X_tensor, y_tensor) loader DataLoader(dataset, batch_size32, shuffleTrue)TensorFlow 的tf.data.Dataset使用链式调用dataset tf.data.Dataset.from_tensor_slices((X, y)) dataset dataset.shuffle(1000).batch(32).prefetch(1)JAX 本身不提供数据加载器实际项目中通常使用tf.data或 NumPy 数组切片。由于 JAX 数组不可变写数据增强时会倾向于“生成新数据”而不是在内存中原地修改X_train_jnp jnp.array(X_train) for start in range(0, len(X_train_jnp), 32): batch X_train_jnp[start:start 32] # 处理当前 batch3.6 训练循环PyTorch 的训练循环是手动展开的你需要自己写zero_grad - forward - backward - stepfor epoch in range(100): optimizer.zero_grad() loss loss_fn(model(X_tensor), y_tensor) loss.backward() optimizer.step()TensorFlow 有两种方式。高层 API 直接用model.fitmodel.compile(optimizeroptimizer, lossloss_fn, metrics[accuracy]) model.fit(X_train, y_train, epochs100, batch_size32)低层 API 则使用GradientTapefor epoch in range(100): with tf.GradientTape() as tape: loss loss_fn(model(X_train), y_train) grads tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))JAX 的训练循环和“面向对象”的 PyTorch 差异最大。每一步更新都是一个函数通常用jax.jit编译加速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 params, opt_state, loss然后在一个普通 Python 循环里不断调用train_step。这里注意首次调用jax.jit函数时会有编译开销之后会快很多。3.7 模型保存、加载与部署PyTorch 通常保存state_dicttorch.save(model.state_dict(), model.pt) model.load_state_dict(torch.load(model.pt))这里要特别提醒PyTorch 2.6 开始torch.load的weights_only参数默认值发生了变化。如果你加载的是旧版本保存的权重并且模型文件中包含自定义类或额外对象可能遇到类似 “Weights only load failed” 的报错这时需要检查weights_only参数设置并且只从可信来源加载模型文件。TensorFlow 保存权重或完整模型model.save_weights(model_weights.h5) model.load_weights(model_weights.h5) # 或者保存完整模型便于部署 model.save(saved_model)JAX 的模型参数是一个普通的 dict你可以用pickle或np.savez保存import pickle with open(params.pkl, wb) as f: pickle.dump(params, f)由于 JAX 参数本身就是普通 Python 结构序列化方式比较自由。但如果要上线部署通常还需要配合 XLA 编译后的产物或者把参数导出给其他推理引擎使用。4. 完整实战三框架实现同一二分类模型4.1 任务说明与数据准备为了真正体现 API 差异这里用一个非常经典的数据集make_moons做二分类任务。数据集有两个特征标签为 0 或 1适合快速验证框架差异。先统一准备数据import numpy as np from sklearn.datasets import make_moons from sklearn.model_selection import train_test_split X, y make_moons(n_samples1000, noise0.1, random_state42) X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42 ) print(X_train.shape, X_test.shape)这段代码三个框架都可以直接使用下面每个框架代码都假设已经执行了上面的数据准备步骤。4.2 PyTorch 实现PyTorch 版本使用nn.Module定义两层 MLP用BCEWithLogitsLoss计算损失torch.optim.Adam更新参数import torch import torch.nn as nn import torch.optim as optim from sklearn.metrics import accuracy_score X_train_t torch.tensor(X_train, dtypetorch.float32) y_train_t torch.tensor(y_train, dtypetorch.float32).reshape(-1, 1) X_test_t torch.tensor(X_test, dtypetorch.float32) y_test_t torch.tensor(y_test, dtypetorch.float32).reshape(-1, 1) class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(2, 16) self.fc2 nn.Linear(16, 1) def forward(self, x): x torch.relu(self.fc1(x)) return self.fc2(x) model MLP() loss_fn nn.BCEWithLogitsLoss() optimizer optim.Adam(model.parameters(), lr0.01) model.train() for epoch in range(200): optimizer.zero_grad() loss loss_fn(model(X_train_t), y_train_t) loss.backward() optimizer.step() if (epoch 1) % 50 0: print(fepoch{epoch 1}, loss{loss.item():.4f}) model.eval() with torch.no_grad(): logits model(X_test_t) preds (torch.sigmoid(logits).numpy() 0.5).astype(int).reshape(-1) print(Test accuracy:, accuracy_score(y_test, preds))PyTorch 的核心流程是“手动展开训练循环”。虽然代码行数多但每一行都在你控制范围内中途打印梯度、保存中间张量都很方便。4.3 TensorFlow 实现TensorFlow 版本使用 Keras 高层 API代码量明显更短import tensorflow as tf from sklearn.metrics import accuracy_score model tf.keras.Sequential([ tf.keras.layers.Dense(16, activationrelu, input_shape(2,)), tf.keras.layers.Dense(1) ]) model.compile( optimizertf.keras.optimizers.Adam(learning_rate0.01), losstf.keras.losses.BinaryCrossentropy(from_logitsTrue), metrics[accuracy] ) history model.fit( X_train, y_train.reshape(-1, 1), epochs200, batch_size32, verbose0 ) print(final training loss:, history.history[loss][-1]) logits model.predict(X_test, verbose0) preds (logits 0.5).astype(int).reshape(-1) print(Test accuracy:, accuracy_score(y_test, preds))Keras 的compile fit把训练循环封装得很完整history对象里还保留了每个 epoch 的 loss 和 accuracy非常方便画曲线。如果你需要细粒度控制可以用 3.6 节里的GradientTape自定义训练循环。4.4 JAX Flax 实现JAX 版本需要同时使用 Flax定义网络和 Optax优化器import jax import jax.numpy as jnp import flax.linen as nn import optax from sklearn.metrics import accuracy_score X_train_j jnp.array(X_train, dtypejnp.float32) y_train_j jnp.array(y_train, dtypejnp.float32).reshape(-1, 1) X_test_j jnp.array(X_test, dtypejnp.float32) y_test_j jnp.array(y_test, dtypejnp.float32).reshape(-1, 1) class MLP(nn.Module): nn.compact def __call__(self, x): x nn.Dense(features16)(x) x nn.relu(x) x nn.Dense(features1)(x) return x model MLP() params model.init(jax.random.PRNGKey(42), X_train_j) optimizer optax.adam(learning_rate0.01) opt_state optimizer.init(params) def loss_fn(params, x, y): logits model.apply(params, x) return optax.sigmoid_binary_cross_entropy(logits, y).mean() 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 params, opt_state, loss for epoch in range(200): params, opt_state, loss train_step(params, opt_state, X_train_j, y_train_j) if (epoch 1) % 50 0: print(fepoch{epoch 1}, loss{loss:.4f}) jax.jit def predict(params, x): return model.apply(params, x) logits predict(params, X_test_j) preds np.array((logits 0.5).astype(int).reshape(-1)) print(Test accuracy:, accuracy_score(y_test, preds))这段代码比较直观地体现了 JAX 的三个特点参数params是显式传入函数的没有隐藏在对象内部计算过程是“纯函数 变换”jax.value_and_grad同时计算 loss 和梯度jax.jit把训练步骤编译成加速版本第一次调用明显变慢之后快速执行。4.5 运行结果对比三个版本在同样数据、同样模型结构、同样优化器参数下运行 200 个 epoch正常情况下最终训练 loss 都会降到很低测试准确率也会接近一致通常可以达到 95% 以上。这里的绝对数值并不重要重要的是你可以直观感受到维度PyTorchTensorFlowJAX代码风格命令式 面向对象高层封装 底层可切换函数式 变换训练循环手动展开fit封装或手动手动写成纯函数调试体验直观可在循环中打断Keras 下较隐蔽函数式需要打印参数首次开销较低较低jit 编译有额外开销5. 常见问题与排查思路5.1 高频报错一览表问题现象常见原因解决思路PyTorch 无法调用 GPUCUDA 版本与 PyTorch 不匹配用torch.cuda.is_available()检查重新按官网命令安装 GPU 版TensorFlow 出现 cuDNN 加载失败CUDA/cuDNN 与 TensorFlow 版本不匹配对照官方版本表安装对应 CUDA 与 cuDNNJAX import 报 libcudart 找不到GPU 版 jaxlib 未安装或 CUDA 路径不对检查pip list中是否有jaxlib确认CUDA_HOME环境变量PyTorch 2.6 加载旧权重报 weights_only 错误torch.load默认权重加载策略变化按实际需求设置weights_only且只加载可信来源文件Keras 训练时 loss 为 NaN学习率过大或标签 shape 不匹配降低学习率检查from_logits设置检查 labels 形状5.2 几个典型问题详解问题 1PyTorch 安装 GPU 版后torch.cuda.is_available()返回 False先确认驱动是否正常运行nvidia-smi。如果显示正常说明驱动没问题那么大概率是 PyTorch 的 CUDA 编译版本与驱动支持的 CUDA 版本不匹配。解决办法是到 PyTorch 官网重新选择对应 CUDA 版本的安装命令不要用默认命令安装到 GPU 机器上。问题 2TensorFlow 安装后报错找不到动态库TensorFlow 对 CUDA 和 cuDNN 的版本要求非常严格。特别是 Linux 下如果系统里存在多个 CUDA 版本很容易出现加载了错误 lib 的情况。建议先用一个干净的 conda 环境根据官方版本表明确安装对应版本的 CUDA toolkit 和 cuDNN不要使用系统里过旧或过新的版本。问题 3JAX 首次运行很慢这是正常现象。jax.jit编译后的函数第一次调用需要做 XLA 编译后续调用才会加速。如果每次运行都慢检查是否每次都在循环内部重新定义了函数或者把jit装饰到了不合适的函数上。一般来说train_step这种需要反复调用的函数才适合jit。6. 最佳实践与选型建议6.1 三类场景选型速查场景推荐框架理由学术研究、论文复现、快速原型PyTorch动态图调试方便Hugging Face 生态兼容好企业服务化部署、移动端/嵌入式场景TensorFlowSavedModel、TF Serving、TFLite 配套完善大规模并行、科学计算、强化学习研究JAXvmap/pmap/jit带来强大的并行和编译能力如果你所在的团队已经有一套成熟的 TensorFlow 训练和部署管线不要因为“PyTorch 论文更多”就盲目迁移成本和风险可能很高。反过来说如果团队从零开始且主要做研究PyTorch 的社区资料和现成代码会帮你省下大量时间。6.2 工程化建议无论选择哪个框架下面几条工程经验都适用锁定版本requirements.txt或pyproject.toml中写清楚框架版本号避免同事之间环境不一致数据管道统一化如果一个团队同时维护多框架代码可以把数据预处理统一放在 NumPy/Pandas 中再分别转成各框架张量减少调试成本模型导出格式前置设计上线前确认需要的部署格式比如 PyTorch 的torch.jit.script、TensorFlow 的 SavedModel、JAX 配合 XLA 的导出方案避免训练完成后才发现部署链路不通梯度检查迁移模型到新框架后先用小数据、固定随机种子跑一个 batch对比梯度数值是否一致日志和可复现性固定随机种子并记录每个关键阶段的 loss、梯度范数、参数范数方便定位训练异常。7. 总结与下一步这篇文章通过 7 个核心 API 维度对比了 PyTorch、TensorFlow、JAX 三大深度学习框架的常见写法差异并用一个二分类实战跑通了三个版本。核心收获有几点PyTorch 是“命令式 对象内部管理状态”适合快速迭代TensorFlow 提供了从高层fit到底层GradientTape的多级抽象生产部署生态完整JAX 是“函数式 显式参数 变换”适合需要编译优化和并行控制的场景框架迁移时最需要关注的是张量操作、自动微分、数据加载和训练循环这 4 个环节因为它们几乎决定了你写代码的整体风格。下一步建议你不要只停留在“看对比”而是亲手把上面三个实战代码各跑一遍再用同一份数据分别输出每个 epoch 的 loss体会三者训练循环的差异。之后可以继续学习 PyTorch 的torch.compile、TensorFlow 的自定义训练循环、JAX 的jax.vmap和jax.pmap这些进阶内容都会在这套 API 基础上展开。如果这篇文章对你有帮助可以先收藏备用后续遇到框架迁移或选型问题再翻出来对照。
返回列表