MirroredStrategy下load_model加载自定义Keras模型多副本失效
问题现象
在MirroredStrategy分布式策略作用域内加载模型时运行不符合预期:已指定4个可见GPU设备,但运行时仅检测到1个计算副本;直接在MirroredStrategy作用域内实例化构造的模型不会出现该问题。
相关背景:
- 复现代码使用继承
tf.keras.models.Model、tf.keras.layers.Layer的自定义子类实现,该自定义实现是异常的核心诱因 - 已验证:在
MirroredStrategy作用域内加载已保存的tf.keras.Sequential模型可正常多副本运行 - 所有日志均来自
TestLayer.call方法中执行tf.print("global_replica_id: {}".format(global_replica_id))的打印输出
复现代码
class Demo(tf.keras.models.Model): def __init__(self, **kwargs): super(Demo, self).__init__(**kwargs) self.test_layer = TestLayer() self.dense_layer = tf.keras.layers.Dense(units=1, activation=None, kernel_initializer="ones", bias_initializer="zeros") def call(self, inputs): vector = self.test_layer(inputs) logit = self.dense_layer(vector) return logit, vector def summary(self): inputs = tf.keras.Input(shape=(10,), dtype=tf.int64) model = tf.keras.models.Model(inputs=inputs, outputs=self.call(inputs)) return model.summary() @tf.function def _step(inputs, labels, model): logit, vector = model(inputs) return logit, vector def tf_dataset(keys, labels, batchsize, repeat): dataset = tf.data.Dataset.from_tensor_slices((keys, labels)) dataset = dataset.repeat(repeat) dataset = dataset.batch(batchsize, drop_remainder=True) return dataset def _dataset_fn(input_context): global_batch_size = 16384 keys = np.ones((global_batch_size, 10)) labels = np.random.randint(low=0, high=2, size=(global_batch_size, 1)) replica_batch_size = input_context.get_per_replica_batch_size(global_batch_size) dataset = tf_dataset(keys, labels, batchsize=replica_batch_size, repeat=1) dataset = dataset.shard(input_context.num_input_pipelines, input_context.input_pipeline_id) return dataset # 在MirroredStrategy作用域内保存模型 strategy = tf.distribute.MirroredStrategy(["GPU:0", "GPU:1", "GPU:2", "GPU:3"]) with strategy.scope(): model = Demo() model.compile() model.summary() dataset = strategy.distribute_datasets_from_function(_dataset_fn) for i, (key_tensors, replica_labels) in enumerate(dataset): print("-" * 30, "step ", str(i), "-" * 30) logit, vector = strategy.run(_step, args=(key_tensors, replica_labels, model)) model.save("demo") # 在MirroredStrategy作用域内加载模型 with strategy.scope(): model2 = tf.keras.models.load_model("demo") dataset = strategy.distribute_datasets_from_function(_dataset_fn) for i, (key_tensors, replica_labels) in enumerate(dataset): print("-" * 30, "step ", str(i), "-" * 30) logit, vector = strategy.run(_step, args=(key_tensors, replica_labels, model2))
运行日志对比
实际运行日志
------------------------------ step 0 ------------------------------ global_replica_id: Tensor("demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:0) global_replica_id: Tensor("replica_1/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:1) global_replica_id: Tensor("replica_2/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:2) global_replica_id: Tensor("replica_3/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:3) 2022-07-13 06:20:56.820402: W tensorflow/python/util/util.cc:368] Sets are not currently considered sequences, but this may change in the future, so consider avoiding using them. ------------------------------ step 0 ------------------------------ global_replica_id: 0
预期运行日志
------------------------------ step 0 ------------------------------ global_replica_id: Tensor("demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:0) global_replica_id: Tensor("replica_1/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:1) global_replica_id: Tensor("replica_2/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:2) global_replica_id: Tensor("replica_3/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:3) 2022-07-13 06:20:56.820402: W tensorflow/python/util/util.cc:368] Sets are not currently considered sequences, but this may change in the future, so consider avoiding using them. ------------------------------ step 0 ------------------------------ global_replica_id: Tensor("demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:0) global_replica_id: Tensor("replica_1/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:1) global_replica_id: Tensor("replica_2/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:2) global_replica_id: Tensor("replica_3/demo/test_layer/replica_id_in_sync_group:0", shape=(), dtype=int32, device=/job:localhost/replica:0/task:0/device:gpu:3)
解决方案
问题根因是自定义Keras子类模型加载时,默认反序列化逻辑不会自动把模型内部变量注册为MirroredStrategy管理的镜像变量,哪怕把load_model写在策略作用域里,也不会触发分布式变量初始化,最终所有计算只会落到单张GPU上。
按以下操作修复:
- 所有自定义的
Layer、Model子类必须实现get_config方法,返回初始化用到的所有自定义参数,保证加载时层结构能被正确重建,不会反序列化出脱离分布式上下文的普通Python对象。示例:
class TestLayer(tf.keras.layers.Layer): # 保留原有__init__、call等方法逻辑 def get_config(self): base_config = super().get_config() # 把__init__方法里自定义的参数按key-value形式加入base_config return base_config class Demo(tf.keras.models.Model): # 保留原有__init__、call、summary等方法逻辑 def get_config(self): base_config = super().get_config() return base_config
- 加载模型时显式传入自定义类的映射字典,加载完成后必须在策略作用域内用和实际输入形状、dtype一致的dummy张量跑一次前向传播,触发分布式镜像变量构建,之后再执行数据集构建、
strategy.run等逻辑:
with strategy.scope(): model2 = tf.keras.models.load_model( "demo", custom_objects={"Demo": Demo, "TestLayer": TestLayer} ) # 构造dummy输入跑一次前向,完成分布式变量初始化 dummy_input = tf.ones((1, 10), dtype=tf.int64) _ = model2(dummy_input)
- 不要用h5格式保存自定义子类模型,保持默认的SavedModel格式即可。h5格式不会保存自定义子类的前向计算图结构,加载时大概率出现上下文绑定错误。
内容的提问来源于stack exchange,提问作者Kingsley Liu
相关产品推荐
相关产品推荐

