You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 14:45:32