mirroredstrategy在单机多卡下不自动生效是因为默认仅使用gpu:0,需显式确认多卡可见、指定devices、将模型定义和compile置于strategy.scope()内,并按单卡调整batch_size。

为什么MirroredStrategy在单机多卡下不自动生效
直接调用 MirroredStrategy 但训练仍只跑在 GPU 0 上,大概率是因为没显式指定可见设备或 TensorFlow 没识别到多张卡。TensorFlow 默认只用 GPU:0,即使你有 4 块卡,MirroredStrategy() 也不会自动接管全部——它只会镜像到当前可见的 GPU 设备列表里。
实操建议:
- 先运行
tf.config.list_physical_devices('GPU')确认是否真能看见多卡(返回长度应 ≥2) - 若只看到 1 张,检查驱动、CUDA 版本与 TF 版本是否匹配,或是否被
CUDA_VISIBLE_DEVICES环境变量限制了可见性 - 显式传入设备列表更稳妥:用
MirroredStrategy(devices=['/gpu:0', '/gpu:1', '/gpu:2', '/gpu:3']),避免依赖默认行为
model.compile() 必须在 strategy.scope() 内完成
把模型定义和 compile() 放在 strategy.scope() 外面,会导致权重不被复制、优化器状态无法同步,最终训练时只有主卡参与计算,其余卡空转甚至报错 ValueError: Variable is not on the current device。
正确写法必须严格遵循作用域包裹:
strategy = tf.distribute.MirroredStrategy()
with strategy.scope():
model = tf.keras.Sequential([...])
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy')
注意:fit() 调用本身不用在 scope 内,但数据输入(如 tf.data.Dataset)建议也提前在 scope 外做好预处理,避免重复构建图。
tf.data.Dataset 的 batch_size 要除以 GPU 数量
很多人直接沿用单卡 batch size,结果每张卡都拿到完整 batch,等效于总 batch size 变成原来的 N 倍,极易 OOM 或梯度爆炸。MirroredStrategy 是“数据并行”,每个设备处理一个子 batch,所以全局 batch size = 单设备 batch size × GPU 数量。
推荐做法:
- 设目标全局 batch size 为 256,4 卡机器就该设
batch_size=64 - 用
dataset.batch(64, drop_remainder=True),别用dataset.batch(256) - 如果用
tf.keras.utils.image_dataset_from_directory等高级 API,其batch_size参数同样要按单卡算
自定义训练循环中 optimizer.apply_gradients() 的位置很关键
在 @tf.function + strategy.run() 的组合里,apply_gradients() 必须放在 strategy.run() 返回的梯度更新逻辑内部,否则会因变量跨设备导致 NotFoundError: Resource xxx does not exist。
典型错误模式是先 strategy.run() 得到梯度,再在 host 上调用 optimizer.apply_gradients();正确结构是:
@tf.function
def train_step(inputs, labels):
with tf.GradientTape() as tape:
predictions = model(inputs, training=True)
loss = loss_fn(labels, predictions)
gradients = tape.gradient(loss, model.trainable_variables)
# ✅ apply_gradients 必须在 strategy.run 内部执行
return strategy.run(optimizer.apply_gradients, args=[(gradients, model.trainable_variables)])
如果你用的是 tf.keras.Model.fit(),这一步由框架自动处理,不用手写;但一旦切到自定义循环,这个细节就绕不开。
设备间同步开销、梯度规约方式、以及 dataset prefetch 深度不够,都会让多卡加速比远低于线性。实际部署前务必用 tf.profiler 抓 trace,看 GPU 利用率是否均衡——经常发现某张卡长期空闲,问题其实出在数据加载瓶颈上,而不是策略本身。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











