外观
深度学习框架
框架选择应服从现有代码、团队经验、目标硬件和生态依赖。论文结果取决于数据、评价和实验控制,不取决于框架名称。
PyTorch
PyTorch 以 Tensor、自动微分、nn.Module 和优化器构成训练闭环:
python
model.train()
for inputs, targets in loader:
inputs, targets = inputs.to(device), targets.to(device)
optimizer.zero_grad(set_to_none=True)
loss = criterion(model(inputs), targets)
loss.backward()
optimizer.step()明确区分 train() 与 eval(),评估时使用无梯度上下文。保存 state_dict、优化器状态、步数、配置和随机数状态,才能真正断点恢复。设备错误先检查 Tensor 的 device、dtype 和 shape;显存不足先缩小 batch、查看未释放计算图,再考虑混合精度和梯度累积。
TensorFlow 与 Keras
Keras 提供高层的 Layer、Model、训练和回调接口,TensorFlow 负责张量、自动微分、数据和部署生态。
python
model = keras.Sequential([
keras.layers.Input((128,)),
keras.layers.Dense(64, activation='relu'),
keras.layers.Dense(10)
])
model.compile(optimizer='adam', loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True))
model.fit(train_ds, validation_data=valid_ds, epochs=10)model.fit 适合标准训练,自定义 train_step 或 GradientTape 适合研究逻辑。使用 tf.data 时检查 shuffle、cache、repeat 的顺序,防止验证数据被打乱或训练循环没有明确步数。
JAX
JAX 用可组合函数变换提供自动微分、JIT、向量化和并行。函数尽量保持纯粹,参数和状态显式传入传出;数组不可变更新和编译期 shape 是常见适应点。
python
import jax
import jax.numpy as jnp
def loss(w, x, y):
return jnp.mean((x @ w - y) ** 2)
grad_fn = jax.jit(jax.grad(loss))
grads = grad_fn(w, x, y)JIT 首次调用包含编译时间,性能比较要区分编译和稳态执行。随机数使用显式 key;修改输入 shape 可能触发重新编译。
PyTorch Lightning
Lightning 把训练步骤、优化器、日志和分布式策略组织进统一结构。它适合标准化训练工程,但论文核心机制仍应清楚落在模块与步骤中,不能被回调和 Trainer 参数遮蔽。
使用前明确:Checkpoint 监控哪个指标、指标是按 step 还是 epoch 聚合、分布式下是否正确同步、恢复是否包含优化器和调度器状态。遇到问题时先用单设备、小数据和最少 Callback 复现,再逐项恢复配置。
如何选择
| 条件 | 倾向 |
|---|---|
| 研究代码和开源模型以 PyTorch 为主 | PyTorch |
| 已有 TensorFlow/Keras 生产或教学生态 | TensorFlow/Keras |
| 需要函数变换、XLA 和可组合并行 | JAX |
| PyTorch 项目训练样板多且团队接受抽象 | Lightning |
跨框架比较必须对齐数据、精度、设备同步和计时边界。任何框架都应先在 5 个样本上跑通输入、前向、损失、反向、保存和加载,再扩大规模。