深度学习框架选型是很多团队迈不过去的第一道坎。同样是训练一个图像分类模型,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 定位差异速查
| 维度 | PyTorch | TensorFlow | JAX |
|---|---|---|---|
| 设计核心 | 命令式动态图 | 高层封装 + 静态图编译 | 纯函数变换 |
| 调试体验 | 可以直接打印中间结果 | Eager 模式也可以,但封装层较多 | 纯函数模式下调试要显式传值 |
| 参数管理 | nn.Module对象统一持有 | Keras 模型内部管理 | 参数是普通数据结构,由用户维护 |
| 推荐场景 | 研究、快速迭代、动态结构模型 | 工业部署、标准化训练管道 | 高性能训练、可微编程、科学计算 |
| 代表生态 | HuggingFace、PyTorch Lightning | Keras、TF Serving、TFLite | Flax、Equinox、Optax、Orbax |
2. 环境准备与安装验证
很多安装问题不是命令写错,而是环境串了。下面先说明如何用 conda 隔离环境,再给出三个框架共存的安装方式,最后给出最小验证脚本。
2.1 用 conda 隔离三套依赖
实际开发时,不建议在同一个 Python 环境里同时安装三个框架。它们对 NumPy、CUDA 版本、protobuf 的要求不完全一致,放在一起容易出现“版本冲突”和“运行时静默替换”的问题。
推荐先建一个独立环境:
conda create -n dl-compare python=3.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 tensorflow==2.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(早期是DeviceArray),API 风格高度对齐 NumPy:
import 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], dtype=torch.float32) x_tf = tf.constant([1, 2], dtype=tf.float32) x_jax = jnp.array([1, 2], dtype=jnp.float32)3.4 张量 API 对比速查
| 场景 | PyTorch | TensorFlow | JAX |
|---|---|---|---|
| 从数据创建 | 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.device | jax.device_put(x, device) |
| 是否可变 | 张量可变,支持 in-place | tf.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:动态计算图 + backward
PyTorch 通过设置requires_grad=True开启梯度追踪:
import torch x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True) 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 TensorFlow:GradientTape 显式记录
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里需要多次求梯度,要设置persistent=True。
4.3 JAX:grad 纯函数变换
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_grad:
value_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 三种微分机制的关键差异
| 维度 | PyTorch | TensorFlow | JAX |
|---|---|---|---|
| 心智模型 | 动态图 + 反向传播 | 显式记录前向轨迹 | 函数到函数的变换 |
| 开启方式 | requires_grad=True | GradientTape上下文 | jax.grad(fn) |
| 控制流 | 原生 Python,完全支持 | Eager 模式支持,tf.function内有限制 | 纯函数内支持,但jit下要处理 trace |
| 非标量 loss | 需要grad_tensors | GradientTape.gradient直接支持 | jax.grad要求输出或归约后是标量 |
| 状态相关 | 参数由模块持有 | 变量由tf.Variable持有 | 状态由用户显式传入函数 |
实际项目中,PyTorch 的“先打印、后 backward”体验最符合直觉,TensorFlow 的“with tape”适合在高层封装外做精细控制,JAX 的函数式求导则要求你把所有输入输出都写清楚,代码结构更规范,但开发时思考成本更高。
5. 核心 API 维度三:模型构建
模型构建 API 决定你如何组织网络结构、管理参数和做初始化。三个框架的差异非常大,迁移时这一层最需要重写。
5.1 PyTorch:nn.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 TensorFlow:Keras 多层封装
TensorFlow 官方推荐使用 Keras。最基础的是Sequential:
import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(32, activation="relu", input_shape=(16,)), tf.keras.layers.Dense(10), ])更复杂的是函数式 API:
inputs = tf.keras.Input(shape=(16,)) x = tf.keras.layers.Dense(32, activation="relu")(inputs) outputs = tf.keras.layers.Dense(10)(x) model = tf.keras.Model(inputs=inputs, outputs=outputs)Keras 模型自带compile、fit、evaluate、save等方法,这些都是 PyTorch 没有的高层封装。优点是把训练管道标准化,缺点是当你要自定义训练逻辑时,要先理解 Keras 封装内部的钩子机制。
5.3 JAX:Flax 模块与参数解耦
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 参数管理与初始化差异
| 维度 | PyTorch | TensorFlow | JAX (Flax) |
|---|---|---|---|
| 模型定义方式 | 继承nn.Module | Sequential/ 函数式 / 子类化 | Flaxnn.Module+@nn.compact |
| 参数对象 | nn.Parameter绑定在模块上 | tf.Variable绑定在层上 | 普通数据结构FrozenDict |
| 参数遍历 | model.parameters() | model.trainable_variables | params字典手动遍历 |
| 初始化方式 | 定义时自动初始化 | 定义时自动初始化 | model.init(key, x)显式初始化 |
| 训练/推理模式 | model.train()/model.eval() | 层参数training=True/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(), lr=1e-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 TensorFlow:model.fit 与自定义训练
TensorFlow 的高层封装把循环细节隐藏起来:
model.compile(optimizer="adam", loss="sparse_categorical_crossentropy") model.fit(train_dataset, epochs=10)如果训练逻辑复杂,可以使用自定义训练循环:
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-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, training=True) 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:显式状态传递与 Optax
JAX 的训练循环没有隐式的模型状态。所有参数、优化器状态、模型内部状态都必须显式传入传出:
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 训练循环核心差异
| 步骤 | PyTorch | TensorFlow | JAX |
|---|---|---|---|
| 梯度清零 | 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 PyTorch:Dataset 与 DataLoader
PyTorch 的数据加载围绕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_size=32, shuffle=True, num_workers=4)DataLoader提供多进程预取、shuffle、batch 分组、collate 函数。这套 API 最大的优点是灵活,你可以完全控制每个样本如何读取和变换。缺点是默认没有数据集缓存,某些场景下需要自己加缓存层。
7.2 TensorFlow:tf.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 项目,需要灵活定义样本读取 | Dataset+DataLoader |
| PyTorch 项目,数据集较小 | 直接 NumPy 切分,减少代码量 |
| TensorFlow 项目,需要高性能输入管道 | tf.data+prefetch |
| TensorFlow 项目,读取 TFRecord | tf.data.TFRecordDataset |
| JAX 项目,数据集较大 | 复用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 TensorFlow:tf.device 与分布式策略
TensorFlow 使用tf.device指定设备:
with tf.device("/GPU:0"): x = tf.random.normal((4, 16))在 Keras 层面,更推荐使用tf.distribute.MirroredStrategy:
strategy = tf.distribute.MirroredStrategy() with strategy.scope(): model = tf.keras.Sequential([...]) model.compile(...) model.fit(train_dataset, epochs=10)分布式策略把设备间的数据分发、梯度同步封装起来,使用fit时非常方便。自定义训练循环时,则需要自己控制strategy.run和strategy.reduce。
8.3 JAX:device_put 与 pmap
JAX 的设备控制主要体现在数据放置和并行变换上:
import jax devices = jax.devices() x = jnp.ones((8, 16)) x_gpu = jax.device_put(x, devices[0])真正的并行能力来自pmap:
@jax.pmap def forward(x): return x * 2pmap会自动把一个 batch 维度切分到多个设备上。这种设计把“数据并行”从手动代码里解放出来,但要求你的函数是纯函数,且输入数据的第一个维度能按设备数整除。
8.4 设备管理差异
| 操作 | PyTorch | TensorFlow | JAX |
|---|---|---|---|
| 指定单设备 | 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/DistributedDataParallel | MirroredStrategy/TPUStrategy | pmap |
| 心智模型 | 显式移动数据 | 设备上下文或策略作用域 | 数据可以放在设备上,函数通过变换并行 |
这里要特别提醒:在 PyTorch 中忘记.to(device)会直接报CUDA error,但 TensorFlow 和 JAX 中 CPU/GPU 切换相对隐式。后两者虽然设备切换更自动,但在混合精度和大 batch 场景下,还是要主动确认数据实际落在哪个设备上,否则性能排查时会很被动。