tf.distribute.MirroredStrategy 是 TensorFlow 2.x 多 GPU 数据并行最稳妥选择,因其自动镜像模型与优化器、支持 NCCL/All-Reduce 同步梯度,并要求严格遵循 strategy.scope()、分布式数据加载、PerReplica 处理及检查点规范。

为什么 tf.distribute.MirroredStrategy 是当前最稳妥的选择
TensorFlow 2.x 中,多 GPU 数据并行训练的官方推荐路径是 tf.distribute.MirroredStrategy,它自动在所有可见 GPU 上复制模型和优化器,并通过 NCCL(Linux)或 NVIDIA Collective Communications Library)同步梯度。如果你用的是 Windows,NCCL 支持有限,MirroredStrategy 会退回到基于 tf.distribute.ReduceOp.SUM 的 All-Reduce 实现,性能略低但依然可用。
常见错误现象包括:训练只跑在 CPU 或单卡上、ValueError: Cannot assign a device for operation...、GPU 显存占用不均衡。这些问题大多源于没正确设置策略作用域或数据输入未适配分布式环境。
- 必须在构建模型前创建并调用
strategy.scope(),否则变量不会被镜像 -
tf.data.Dataset必须用strategy.experimental_distribute_dataset()包装,或直接在strategy.run()内部做batch和prefetch - 避免在
strategy.scope()外定义model或optimizer,否则它们不会被分布到各 GPU
如何正确包装模型、损失和训练步骤
核心原则是:所有可训练变量、前向/反向计算、梯度更新都必须落在 strategy.scope() 内;而数据加载、评估逻辑、保存检查点可以放在外面。
一个典型结构是把训练逻辑封装成带 @tf.function 装饰的函数,再传给 strategy.run()。注意:不能直接对 strategy.run() 返回值做 Python 操作(比如取均值、打印),它返回的是 PerReplica 对象,需用 strategy.reduce() 合并。
立即学习“Python免费学习笔记(深入)”;
- 损失函数应返回 per-example loss(而非 batch mean),让
strategy.reduce()统一做加权平均 - 使用
tf.keras.losses.sparse_categorical_crossentropy(..., reduction=tf.keras.losses.Reduction.NONE)避免内部提前求均 - 梯度裁剪应在
strategy.run()内完成,否则各卡梯度未合并,裁剪失效
@tf.function
def train_step(inputs):
images, labels = inputs
with tf.GradientTape() as tape:
predictions = model(images, training=True)
loss = loss_fn(labels, predictions)
gradients = tape.gradient(loss, model.trainable_variables)
gradients = [tf.clip_by_norm(g, 1.0) for g in gradients]
optimizer.apply_gradients(zip(gradients, model.trainable_variables))
return loss
<h1>在 strategy.run 中调用</h1><p>per_replica_losses = strategy.run(train_step, args=(per_replica_inputs,))
loss = strategy.reduce(tf.distribute.ReduceOp.MEAN, per_replica_losses, axis=None)
数据加载必须适配 MirroredStrategy 分布模式
最常被忽略的一点:tf.data.Dataset 默认不感知 GPU 分布。若直接用 dataset.batch(batch_size) 并传入 strategy.experimental_distribute_dataset(),实际每个 GPU 会拿到完整 batch 的子集(即 global_batch_size = per_gpu_batch_size × num_gpus),但你得自己确保 batch_size 是 GPU 数量的整数倍,否则最后一批会出错。
更健壮的做法是先设好 per-GPU batch size,再用 global_batch_size = per_gpu_batch_size * strategy.num_replicas_in_sync 构建 dataset,并启用 drop_remainder=True。
- 务必调用
dataset = dataset.cache().shuffle(buffer_size).repeat()在distributed_dataset包装前完成,否则 shuffle 和 repeat 会在每张卡上独立执行,打乱全局顺序 -
prefetch(tf.data.AUTOTUNE)应放在distributed_dataset之后,否则 prefetch 无法跨设备生效 - 避免在
map()函数中使用非 tf ops(如 PIL、cv2),它们不支持分布式图执行
检查点保存与恢复的陷阱
tf.train.Checkpoint 本身支持分布式变量,但保存路径必须由主进程(host device)统一写入,且恢复时要确保所有 GPU 变量已初始化完毕。常见错误是:恢复后某张卡上的权重仍是随机初始化值,导致训练崩溃。
关键动作是:在 strategy.scope() 内初始化模型和优化器后,再构建 Checkpoint;保存时用 checkpoint.save(),恢复时用 checkpoint.restore().expect_partial()(避免因新增/删减变量报错)。
- 不要在
strategy.run()内调用save(),会导致多进程重复写文件 - 如果用了
tf.keras.Model.save_weights(),必须指定save_format='h5'或'tf',且路径需为本地文件系统路径(不支持 GCS/S3 直接写) - 验证是否恢复成功:可在恢复后打印
model.variables[0].numpy()[0],确认各卡值一致
多 GPU 训练真正难的不是启动,而是让每张卡上的计算图、数据流、状态更新完全对齐。哪怕只漏掉一个 strategy.scope() 或一处 drop_remainder,都可能让训练悄无声息地降级成单卡运行。


















