拓冰建站拓冰建站
首页 / 资讯中心 / 正文

深度学习框架选型:PyTorch、TensorFlow、JAX七个核心API维度对比

深度学习框架选型是很多团队迈不过去的第一道坎。同样是训练一个图像分类模型PyTorch、TensorFlow、JAX 写出来的代码结构差异非常大这种差异并不是语法习惯不同而是框架对“计算图构建、自动微分、参数管理、设备调度”这四件事采用了完全不同的设计立场。标题里的“7”指的不是 7 个框架而是 7 个核心 API 维度。本文从张量创建与基础运算、自动微分、模型构建、训练循环、数据加载与预处理、设备管理与并行、模型保存与部署这 7 个层面横向对比 PyTorch、TensorFlow、JAX 三个框架。适合已经会用其中一个框架、想快速迁移到另外两个的开发者也适合准备选型但不知道从何下手的初学者。读完以后你能在遇到框架差异、代码迁移、版本兼容和训练循环改造时快速判断问题出在框架设计层面还是自己的代码层面。1. 三个框架的定位差异与 API 设计哲学看 API 之前先理解三个框架各自服务谁。API 只是外在表达背后是框架作者对“深度学习开发哪里最重要”这个问题的回答。1.1 PyTorch命令式动态图Python 原生体验PyTorch 的核心设计是命令式动态图。用户写下的每一行张量运算都会在真实执行过程中被 autograd 自动追踪。调用loss.backward()时梯度会沿着实际执行路径回传。这种设计带来一个非常大的好处调试直观。模型中间层的输出可以用print直接打印可以在forward方法里加断点可以用 Python 的if、for来写控制流不需要把逻辑改造成框架规定的语法。对于研究型项目、论文复现和快速原型验证这种体验几乎没有替代品。代价也有动态图在极端性能场景下优化空间不如静态图大。PyTorch 后面的torch.compile、TorchScript 都是在弥补这个短板但默认心智模型仍然是“Python 跑到哪里就算到哪里”。1.2 TensorFlow从静态图走向 Keras 封装面向生产部署TensorFlow 1.x 时代的主推模式是静态图先定义tf.Graph再用tf.Session执行。这种模式在服务端部署上有优势但调试体验很差社区一度吐槽很多。TensorFlow 2.x 做了两个关键调整默认启用 Eager Execution动态执行同时把 Keras 提升为官方高层 API。现在写 TensorFlow 时多数入口是tf.keras而不是底层的tf.Graph。tf.function可以把你写的 Python 函数编译成计算图既保留开发阶段的灵活性又能在部署阶段获得静态图性能。因此TensorFlow 的 API 有一种“高层封装优先”的特点。model.fit、model.compile、model.evaluate这些方法把训练循环高度封装起来适合快速起步也适合做标准化的生产管道。缺点是当你的训练逻辑比较特殊时要绕过封装去改底层学习曲线会明显变陡。1.3 JAX函数式变换把 NumPy 变成可微编程语言JAX 的定位不是“有了 PyTorch 为什么还要 JAX”的替代品而是对“数值计算 自动微分 硬件加速”的一次重新抽象。它的底层假设是一切都是纯函数变换。同一个函数可以被jax.jit编译加速被jax.grad求导被jax.vmap自动向量化被jax.pmap分布到多设备。由于函数没有外部副作用所有状态都显式传入传出JAX 可以放心地对函数做组合和变换。这个设计很优雅但也带来学习成本。JAX 没有 PyTorch 那种“模型对象 参数对象”的类式管理需要你自己管理参数结构。实际项目中通常会借助 Flax、Equinox、Optax 这类生态库来降低使用门槛但核心 API 仍然是函数式的。1.4 定位差异速查维度PyTorchTensorFlowJAX设计核心命令式动态图高层封装 静态图编译纯函数变换调试体验可以直接打印中间结果Eager 模式也可以但封装层较多纯函数模式下调试要显式传值参数管理nn.Module对象统一持有Keras 模型内部管理参数是普通数据结构由用户维护推荐场景研究、快速迭代、动态结构模型工业部署、标准化训练管道高性能训练、可微编程、科学计算代表生态HuggingFace、PyTorch LightningKeras、TF Serving、TFLiteFlax、Equinox、Optax、Orbax2. 环境准备与安装验证很多安装问题不是命令写错而是环境串了。下面先说明如何用 conda 隔离环境再给出三个框架共存的安装方式最后给出最小验证脚本。2.1 用 conda 隔离三套依赖实际开发时不建议在同一个 Python 环境里同时安装三个框架。它们对 NumPy、CUDA 版本、protobuf 的要求不完全一致放在一起容易出现“版本冲突”和“运行时静默替换”的问题。推荐先建一个独立环境conda create -n dl-compare python3.10 -y conda activate dl-comparePython 版本建议选择 3.10 或 3.11。TensorFlow 对 Python 版本要求较严格JAX 在部分版本上适配较慢选 3.10 在三个框架之间兼容性最好。如果你的系统是 Windows还要注意 JAX 对 Windows 的原生支持有限建议使用 WSL2 或 Linux 环境。2.2 安装命令与版本匹配安装 PyTorch 时官方推荐从官网生成命令。下面是以 CUDA 12.1 为例的安装方式pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121安装 TensorFlow 时注意看版本号。例如 TensorFlow 2.18 对 Python 3.9 到 3.12 支持较好但不一定兼容所有操作系统安装前要确认自己的 Python 版本pip install tensorflow2.18安装 JAX 时Linux 环境可以直接安装带 CUDA 支持的版本pip install jax jaxlib flax optax如果安装 jaxlib 后无法使用 GPU通常是因为 pip 默认安装了 CPU 版本。这时要按 JAX 官方文档的指引安装与本地 CUDA 版本对应的 jaxlib。2.3 安装后验证 GPU 和版本安装完成后不要急着写模型。先运行一段最小验证脚本确认三个框架都能正常导入并且 GPU 可用import torch import tensorflow as tf import jax import jax.numpy as jnp print(PyTorch:, torch.__version__, CUDA:, torch.cuda.is_available()) print(TensorFlow:, tf.__version__, GPU:, tf.config.list_physical_devices(GPU)) print(JAX:, jax.__version__, devices:, jax.devices())预期输出中PyTorch 和 TensorFlow 显示 GPU 设备JAX 能打印出cuda:0或gpu:0设备。如果某个框架没有打印出 GPU先不要继续安装其他依赖优先解决该框架的 GPU 配置问题否则后面跑模型时会浪费大量排错时间。注意验证脚本通过并不代表后续所有操作都会正常。不要在torch.__version__正常打印后忽略torch.cuda.is_available()为 False 的情况这是环境未配好最常见的信号。3. 核心 API 维度一张量创建与基础运算张量是三个框架最底层的 API也是初学者最容易混淆的地方。它们的核心区别是是否可变、是否显式管理设备、默认有哪些便捷构造函数。3.1 三种张量的创建方式PyTorch 使用torch.Tensor创建方式非常直接import torch x torch.randn(4, 16) # 标准正态分布 y torch.zeros(4, 16) # 全零 z torch.tensor([[1.0, 2.0]]) # 从数据创建TensorFlow 使用tf.Tensor创建语法与 NumPy 接近import tensorflow as tf x tf.random.normal((4, 16)) y tf.zeros((4, 16)) z tf.constant([[1.0, 2.0]])JAX 使用jax.Array早期是DeviceArrayAPI 风格高度对齐 NumPyimport jax.numpy as jnp x jnp.ones((4, 16)) y jnp.zeros((4, 16)) z jnp.asarray([[1.0, 2.0]])从创建方式上看三者差异不大。真正的差异在“可变性”。3.2 可变性与 in-place 操作PyTorch 的张量默认是可变的支持x.add_(1)这类 in-place 操作。这在内存优化时很有用但也容易引入 bug因为 autograd 对 in-place 操作追踪有限制。TensorFlow 的tf.Tensor不可变每次运算都产生新张量。如果要保存可变参数需要显式使用tf.Variable。JAX 的数组也不可变。y x 1会返回新数组原来的x不会改变。这是 JAX 纯函数设计的基础也是新手最容易踩坑的地方习惯性地以为x 1修改了原变量实际上你需要重新赋值。3.3 数据类型与默认值三个框架默认都是float32但细节不同。PyTorch 的torch.tensor([1, 2])会推导出整数类型而torch.randn默认是float32。TensorFlow 的tf.constant([1, 2])推导为int32。JAX 的jnp.array([1, 2])也是整数类型。在数据集准备阶段推荐显式指定 dtype避免因为默认类型不一致导致计算异常x_torch torch.tensor([1, 2], dtypetorch.float32) x_tf tf.constant([1, 2], dtypetf.float32) x_jax jnp.array([1, 2], dtypejnp.float32)3.4 张量 API 对比速查场景PyTorchTensorFlowJAX从数据创建torch.tensor(...)tf.constant(...)jnp.array(...)正态分布随机torch.randn(size)tf.random.normal(shape)jax.random.normal(key, shape)全零torch.zeros(size)tf.zeros(shape)jnp.zeros(shape)设备转移x.to(cuda)tf.identity(x)tf.devicejax.device_put(x, device)是否可变张量可变支持 in-placetf.Tensor不可变tf.Variable可变数组不可变Python 控制流原生支持原生支持但tf.function内有约束原生支持但jax.jit内需要 trace 兼容注意 JAX 的随机数生成和其他两个框架有本质区别。jax.random.normal(key, shape)必须显式传入PRNGKey这是因为 JAX 坚持纯函数的无副作用原则。写 JAX 随机数时不能像 PyTorch 那样依赖全局随机种子否则在jit或vmap里会得到难以排查的随机数重复问题。4. 核心 API 维度二自动微分自动微分是深度学习框架的核心三个框架在这一层的 API 设计差异最大。理解这一层就能理解为什么相同训练逻辑在三个框架里写出来完全不同。4.1 PyTorch动态计算图 backwardPyTorch 通过设置requires_gradTrue开启梯度追踪import torch x torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) y (x ** 2).sum() y.backward() print(x.grad) # tensor([2., 4., 6.])执行y.backward()时PyTorch 会从y开始沿着反向计算图一路传播梯度到所有叶子张量。这里的关键点是计算图是在前向执行过程中动态构建的所以每一层print、每个if分支都会真实反映在执行路径里。使用上要注意只有标量才能直接调用backward()。如果loss不是标量需要传入grad_tensors参数或者先做sum()、mean()等归约。4.2 TensorFlowGradientTape 显式记录TensorFlow 使用tf.GradientTape来记录前向轨迹import tensorflow as tf x tf.Variable([1.0, 2.0, 3.0]) with tf.GradientTape() as tape: y tf.reduce_sum(x ** 2) grad tape.gradient(y, x) print(grad.numpy()) # [2. 4. 6.]GradientTape默认只记录tf.Variable。如果你要计算普通张量的梯度需要显式tape.watch(x)。每个tape.gradient调用只能执行一次因为默认的梯度记录会在调用后释放资源。如果同一个tape里需要多次求梯度要设置persistentTrue。4.3 JAXgrad 纯函数变换JAX 的思路完全不同。它不记录任何轨迹而是把求导当成一个函数变换import jax import jax.numpy as jnp def loss_fn(x): return jnp.sum(x ** 2) grad_fn jax.grad(loss_fn) print(grad_fn(jnp.array([1.0, 2.0, 3.0]))) # [2. 4. 6.]jax.grad接收一个函数返回一个导函数。这个导函数会以前向函数完全相同的输入参数作为输入返回梯度。如果想要同时得到损失值和梯度使用jax.value_and_gradvalue_and_grad_fn jax.value_and_grad(loss_fn) loss, grads value_and_grad_fn(jnp.array([1.0, 2.0, 3.0]))这里的关键约束是jax.grad只能作用于纯函数。函数内部不能修改外部全局状态不能依赖随机数全局种子不能有print这样的副作用。否则在jit编译后会产生不符合预期的行为。4.4 三种微分机制的关键差异维度PyTorchTensorFlowJAX心智模型动态图 反向传播显式记录前向轨迹函数到函数的变换开启方式requires_gradTrueGradientTape上下文jax.grad(fn)控制流原生 Python完全支持Eager 模式支持tf.function内有限制纯函数内支持但jit下要处理 trace非标量 loss需要grad_tensorsGradientTape.gradient直接支持jax.grad要求输出或归约后是标量状态相关参数由模块持有变量由tf.Variable持有状态由用户显式传入函数实际项目中PyTorch 的“先打印、后 backward”体验最符合直觉TensorFlow 的“with tape”适合在高层封装外做精细控制JAX 的函数式求导则要求你把所有输入输出都写清楚代码结构更规范但开发时思考成本更高。5. 核心 API 维度三模型构建模型构建 API 决定你如何组织网络结构、管理参数和做初始化。三个框架的差异非常大迁移时这一层最需要重写。5.1 PyTorchnn.Module 对象化管理PyTorch 通过继承nn.Module来定义模型import torch.nn as nn class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(16, 32) self.relu nn.ReLU() self.fc2 nn.Linear(32, 10) def forward(self, x): return self.fc2(self.relu(self.fc1(x))) model MLP()参数全部由模块内部持有通过model.parameters()可以遍历所有参数通过model.state_dict()可以拿到参数和 buffer 的字典。训练前用model.train()切换到训练模式用model.eval()切换到推理模式。5.2 TensorFlowKeras 多层封装TensorFlow 官方推荐使用 Keras。最基础的是Sequentialimport tensorflow as tf model tf.keras.Sequential([ tf.keras.layers.Dense(32, activationrelu, input_shape(16,)), tf.keras.layers.Dense(10), ])更复杂的是函数式 APIinputs tf.keras.Input(shape(16,)) x tf.keras.layers.Dense(32, activationrelu)(inputs) outputs tf.keras.layers.Dense(10)(x) model tf.keras.Model(inputsinputs, outputsoutputs)Keras 模型自带compile、fit、evaluate、save等方法这些都是 PyTorch 没有的高层封装。优点是把训练管道标准化缺点是当你要自定义训练逻辑时要先理解 Keras 封装内部的钩子机制。5.3 JAXFlax 模块与参数解耦JAX 本身没有模型对象实际项目通常使用 Flax。Flax 的nn.Module更像参数构造函数而不是运行时对象from flax import linen as nn import jax import jax.numpy as jnp class MLP(nn.Module): features: int 32 nn.compact def __call__(self, x): x nn.Dense(self.features)(x) x nn.relu(x) x nn.Dense(10)(x) return x model MLP() params model.init(jax.random.PRNGKey(0), jnp.ones((1, 16))) y model.apply(params, jnp.ones((1, 16)))注意这里的关键点params是一个普通的数据结构FrozenDict模型本身不保存参数。调用model.apply(params, x)时才把参数显式传进去。训练时你需要自己把params传给损失函数再传给value_and_grad。5.4 参数管理与初始化差异维度PyTorchTensorFlowJAX (Flax)模型定义方式继承nn.ModuleSequential/ 函数式 / 子类化Flaxnn.Modulenn.compact参数对象nn.Parameter绑定在模块上tf.Variable绑定在层上普通数据结构FrozenDict参数遍历model.parameters()model.trainable_variablesparams字典手动遍历初始化方式定义时自动初始化定义时自动初始化model.init(key, x)显式初始化训练/推理模式model.train()/model.eval()层参数trainingTrue/False需要自己在函数里控制 batchnorm/dropout 状态常见误区是“JAX 也有 nn.Module是不是用法和 PyTorch 一样”。实际上 Flax 的模块只在init时负责生成参数真正推理和训练时要把参数传回apply。如果你带着 PyTorch 的思维去写很容易把网络结构里的小模块实例保存下来反复调用结果发现参数没有更新。6. 核心 API 维度四训练循环训练循环是三个框架使用体验差距最明显的地方。PyTorch 几乎完全手写TensorFlow 有fit封装JAX 则要求显式管理所有状态。6.1 PyTorch手动梯度清零、反向传播、参数更新import torch import torch.nn as nn optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.CrossEntropyLoss() for epoch in range(10): for x, y in dataloader: optimizer.zero_grad() logits model(x) loss loss_fn(logits, y) loss.backward() optimizer.step()PyTorch 把训练循环完全暴露给开发者。zero_grad负责清空梯度backward负责计算梯度step负责更新参数。好处是每一个环节都清晰可控坏处是初学者容易忘记zero_grad()导致梯度累加。6.2 TensorFlowmodel.fit 与自定义训练TensorFlow 的高层封装把循环细节隐藏起来model.compile(optimizeradam, losssparse_categorical_crossentropy) model.fit(train_dataset, epochs10)如果训练逻辑复杂可以使用自定义训练循环optimizer tf.keras.optimizers.Adam(learning_rate1e-3) loss_fn tf.keras.losses.SparseCategoricalCrossentropy() for epoch in range(10): for x, y in train_dataset: with tf.GradientTape() as tape: logits model(x, trainingTrue) loss loss_fn(y, logits) grads tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))这里可以看到 TensorFlow 的自定义训练循环和 PyTorch 结构类似但没有zero_grad因为GradientTape每次重新进入会自动重建记录器。6.3 JAX显式状态传递与 OptaxJAX 的训练循环没有隐式的模型状态。所有参数、优化器状态、模型内部状态都必须显式传入传出import jax import optax optimizer optax.adam(1e-3) opt_state optimizer.init(params) jax.jit def train_step(params, opt_state, x, y): def loss_fn(params): logits model.apply(params, x) return jnp.mean(optax.softmax_cross_entropy_with_integer_labels(logits, y)) loss, grads jax.value_and_grad(loss_fn)(params) updates, opt_state optimizer.update(grads, opt_state, params) params optax.apply_updates(params, updates) return params, opt_state, loss for epoch in range(10): for x, y in dataset: params, opt_state, loss train_step(params, opt_state, x, y)这段代码里最值得注意的点是jax.jit。被 jit 编译后train_step内部的 Python 循环和控制流都会被跟踪优化但你无法在函数内部print中间张量或依赖外部全局状态。调试阶段可以先去掉jax.jit确认逻辑正确后再开启编译。6.4 训练循环核心差异步骤PyTorchTensorFlowJAX梯度清零optimizer.zero_grad()不需要无状态不需要计算梯度loss.backward()tape.gradient(loss, vars)jax.value_and_grad(fn)更新参数optimizer.step()optimizer.apply_gradients(...)optax.apply_updates(...)历史梯度累积会累积需手动清零GradientTape 重建不累积纯函数无累积随机状态全局随机种子全局随机种子显式传入 PRNGKey如果只从训练循环的代码量看TensorFlow 的fit最短JAX 最长。但长度不代表优劣。JAX 把每个状态都写清楚后多卡并行时反而更好推断因为每个设备上的状态变化都是显式的。7. 核心 API 维度五数据加载与预处理数据管道是真实项目中很容易被低估的一环。三个框架在这一层的 API 设计思路差异很大。7.1 PyTorchDataset 与 DataLoaderPyTorch 的数据加载围绕Dataset和DataLoader两个类from torch.utils.data import Dataset, DataLoader class MyDataset(Dataset): def __init__(self, x, y): self.x x self.y y def __len__(self): return len(self.x) def __getitem__(self, idx): return self.x[idx], self.y[idx] dataloader DataLoader(MyDataset(x, y), batch_size32, shuffleTrue, num_workers4)DataLoader提供多进程预取、shuffle、batch 分组、collate 函数。这套 API 最大的优点是灵活你可以完全控制每个样本如何读取和变换。缺点是默认没有数据集缓存某些场景下需要自己加缓存层。7.2 TensorFlowtf.data 流水线TensorFlow 使用tf.data.Dataset构建数据管道dataset tf.data.Dataset.from_tensor_slices((x, y)) dataset dataset.shuffle(1000).batch(32).prefetch(1)tf.data把“读取-变换-缓冲-预取”设计成流水线算子适合构建高性能输入管道。prefetch(1)让数据加载与模型训练并行执行能显著减少 GPU 等待。在 TensorFlow 中数据预处理尽量放在dataset.map里而不是模型内部这样可以充分利用流水线并行。7.3 JAX没有官方数据加载器JAX 没有自己的DataLoader或tf.data。常见做法有两个第一种直接用 NumPy 数组和手动批次切分for i in range(0, len(x), batch_size): x_batch x[i:i batch_size] y_batch y[i:i batch_size] params, opt_state, loss train_step(params, opt_state, x_batch, y_batch)第二种复用tf.data作为数据管道import tensorflow as tf dataset tf.data.Dataset.from_tensor_slices((x, y)).batch(32).prefetch(1) for x_batch, y_batch in dataset: x_batch jnp.asarray(x_batch.numpy()) y_batch jnp.asarray(y_batch.numpy()) params, opt_state, loss train_step(params, opt_state, x_batch, y_batch)JAX 社区没有把数据加载做成核心 API原因很简单数据加载本质上是 I/O 和预处理问题与自动微分和编译无关。把数据管道交给成熟的tf.data或外部库是更务实的做法。7.4 数据管道选型建议需求推荐方案PyTorch 项目需要灵活定义样本读取DatasetDataLoaderPyTorch 项目数据集较小直接 NumPy 切分减少代码量TensorFlow 项目需要高性能输入管道tf.dataprefetchTensorFlow 项目读取 TFRecordtf.data.TFRecordDatasetJAX 项目数据集较大复用tf.data或datasets库JAX 项目必须保持纯函数风格NumPy 切片 jnp.asarray实际项目中常见错误是把tf.data.Dataset直接传给 PyTorch 模型或者把DataLoader返回的torch.Tensor直接传给 JAX 的jnp运算。三个框架的数据类型并不隐式兼容跨框架传递前必须显式转换。8. 核心 API 维度六设备管理与并行设备管理是深度学习工程化的基础。三个框架的设备 API 风格完全不同迁移代码时最容易在这层踩坑。8.1 PyTorch.to(device) 移动一切PyTorch 使用torch.device统一表示设备和设备类型device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) x x.to(device)PyTorch 的to方法会递归移动模块内的所有参数和 buffer。多卡训练可以使用torch.nn.DataParallel大规模训练通常使用torch.distributed和DistributedDataParallel。8.2 TensorFlowtf.device 与分布式策略TensorFlow 使用tf.device指定设备with tf.device(/GPU:0): x tf.random.normal((4, 16))在 Keras 层面更推荐使用tf.distribute.MirroredStrategystrategy tf.distribute.MirroredStrategy() with strategy.scope(): model tf.keras.Sequential([...]) model.compile(...) model.fit(train_dataset, epochs10)分布式策略把设备间的数据分发、梯度同步封装起来使用fit时非常方便。自定义训练循环时则需要自己控制strategy.run和strategy.reduce。8.3 JAXdevice_put 与 pmapJAX 的设备控制主要体现在数据放置和并行变换上import jax devices jax.devices() x jnp.ones((8, 16)) x_gpu jax.device_put(x, devices[0])真正的并行能力来自pmapjax.pmap def forward(x): return x * 2pmap会自动把一个 batch 维度切分到多个设备上。这种设计把“数据并行”从手动代码里解放出来但要求你的函数是纯函数且输入数据的第一个维度能按设备数整除。8.4 设备管理差异操作PyTorchTensorFlowJAX指定单设备x.to(cuda)tf.device(/GPU:0)jax.device_put(x, device)判断 GPU 可用torch.cuda.is_available()tf.config.list_physical_devices(GPU)jax.devices()多卡训练DataParallel/DistributedDataParallelMirroredStrategy/TPUStrategypmap心智模型显式移动数据设备上下文或策略作用域数据可以放在设备上函数通过变换并行这里要特别提醒在 PyTorch 中忘记.to(device)会直接报CUDA error但 TensorFlow 和 JAX 中 CPU/GPU 切换相对隐式。后两者虽然设备切换更自动但在混合精度和大 batch 场景下还是要主动确认数据实际落在哪个设备上否则性能排查时会很被动。9. 核心 API
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门