如何替代TensorFlow2.4已移除的keras.utils.multi_gpu_model实现多GPU训练
TensorFlow 2.4+ 多GPU训练替代方案
tensorflow.keras.utils.multi_gpu_model接口从TensorFlow 2.4版本开始正式移除,官方推荐使用tf.distribute.MirroredStrategy作为单主机多GPU场景的替代方案,实现逻辑比原有接口更简洁,训练性能也更优。
适配后的代码实现
完全对齐原有代码的功能:加载本地模型、自动适配多GPU环境,输出的模型可直接用于训练/预测,原有后续调用逻辑无需修改:
import tensorflow as tf from tensorflow.keras.models import load_model # 统计可用GPU数量 gpus = len(tf.config.list_physical_devices('GPU')) if gpus > 1: # 初始化分布式策略,默认自动识别所有可用GPU strategy = tf.distribute.MirroredStrategy() # 在策略上下文内加载模型,自动完成多GPU适配 with strategy.scope(): model = load_model("my_model.h5") else: # 单GPU/CPU场景走原有逻辑 model = load_model("my_model.h5")
注意事项
- 可以在初始化
MirroredStrategy时手动指定要使用的GPU,比如MirroredStrategy(devices=["/gpu:0", "/gpu:1"])表示仅使用前两块GPU - 适配后的模型调用
compile()、fit()、predict()的逻辑和单卡场景完全一致,不需要额外修改其他代码 - 自定义训练循环场景下,仅需要把训练步骤用
strategy.run()包裹即可,常规建模场景无需额外改动
内容的提问来源于stack exchange,提问作者matt
相关产品推荐
相关产品推荐

