深度学习框架选型:PyTorch、TensorFlow、JAX七个核心API维度对比
2026/8/31 15:32:10 网站建设 项目流程

深度学习框架选型是很多团队迈不过去的第一道坎。同样是训练一个图像分类模型,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 的iffor来写控制流,不需要把逻辑改造成框架规定的语法。对于研究型项目、论文复现和快速原型验证,这种体验几乎没有替代品。

代价也有:动态图在极端性能场景下,优化空间不如静态图大。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.Graphtf.function可以把你写的 Python 函数编译成计算图,既保留开发阶段的灵活性,又能在部署阶段获得静态图性能。

因此,TensorFlow 的 API 有一种“高层封装优先”的特点。model.fitmodel.compilemodel.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、Orbax

2. 环境准备与安装验证

很多安装问题不是命令写错,而是环境串了。下面先说明如何用 conda 隔离环境,再给出三个框架共存的安装方式,最后给出最小验证脚本。

2.1 用 conda 隔离三套依赖

实际开发时,不建议在同一个 Python 环境里同时安装三个框架。它们对 NumPy、CUDA 版本、protobuf 的要求不完全一致,放在一起容易出现“版本冲突”和“运行时静默替换”的问题。

推荐先建一个独立环境:

conda create -n dl-compare python=3.10 -y conda activate dl-compare

Python 版本建议选择 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:0gpu: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 对比速查

场景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 那样依赖全局随机种子,否则在jitvmap里会得到难以排查的随机数重复问题。

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 三种微分机制的关键差异

维度PyTorchTensorFlowJAX
心智模型动态图 + 反向传播显式记录前向轨迹函数到函数的变换
开启方式requires_grad=TrueGradientTape上下文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 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 模型自带compilefitevaluatesave等方法,这些都是 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 参数管理与初始化差异

维度PyTorchTensorFlowJAX (Flax)
模型定义方式继承nn.ModuleSequential/ 函数式 / 子类化Flaxnn.Module+@nn.compact
参数对象nn.Parameter绑定在模块上tf.Variable绑定在层上普通数据结构FrozenDict
参数遍历model.parameters()model.trainable_variablesparams字典手动遍历
初始化方式定义时自动初始化定义时自动初始化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 训练循环核心差异

步骤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 PyTorch:Dataset 与 DataLoader

PyTorch 的数据加载围绕DatasetDataLoader两个类:

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 没有自己的DataLoadertf.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 项目,读取 TFRecordtf.data.TFRecordDataset
JAX 项目,数据集较大复用tf.datadatasets
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.distributedDistributedDataParallel

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.runstrategy.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 * 2

pmap会自动把一个 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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询