tf.keras加载模型遇MeanMetricWrapper恢复失败求助
问题详情
花费大量时间训练的Keras模型,通过tf.keras.models.save_model保存后无法正常加载,报错如下:
ValueError: Unable to restore custom object of class "MeanMetricWrapper" (type _tf_keras_metric). Please make sure that this class is included in the
custom_objectsarg when callingload_model(). Also, check that the class implementsget_configandfrom_config.Complete metadata: {'class_name': 'MeanMetricWrapper', 'name': 'loss', 'dtype': 'float32', 'config': {'name': 'loss', 'dtype': 'float32'}, 'shared_object_id': 12}
仅使用了继承自keras.Model的自定义模型类CustomModel,重写了__init__和train_step方法,未使用自定义指标。自定义__init__代码:
def __init__(self, *args, **kwargs): self.pool = kwargs["pool"] del kwargs["pool"] super().__init__(*args, **kwargs)
该方法用于传入多进程Pool供train_step使用。
模型训练与保存代码:
model = CustomModel(inputs=inputs, outputs=outputs, pool=pool) model.compile(optimizer="adam", metrics=["loss"]) model.fit([0]*1024*32, [0]*1024*32, epochs=20, batch_size=1024) tf.keras.models.save_model(model, "C:\src\models\model")
(自定义train_step未实际使用输入的x和y数据)
加载尝试及对应报错
- 指定
MeanMetricWrapper加载:
tf.keras.models.load_model("C:\src\models\model", custom_objects={"MeanMetricWrapper": tf.metrics.MeanMetricWrapper})
仍出现相同的MeanMetricWrapper错误。
- 指定
CustomModel加载:
tf.keras.models.load_model("C:\src\models\model", custom_objects={"CustomModel": CustomModel})
出现新错误:
File "C:\Users...\venv\Lib\site-packages\tensorflow\python\trackable\base.py", line 204, in _method_wrapper
result = method(self, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
TypeError: Functional.init() missing 2 required positional arguments: 'inputs' and 'outputs'
运行环境:Windows 10、PyCharm、pip安装的tensorflow-intel。
修复方案
方案1:完善CustomModel的序列化逻辑
问题核心是自定义模型CustomModel的__init__修改了参数,导致序列化时无法正确重建。需要为CustomModel实现get_config和from_config方法(注意:多进程Pool无法序列化,需加载后手动传入):
class CustomModel(tf.keras.Model): def __init__(self, *args, **kwargs): self.pool = kwargs.pop("pool", None) # 用pop避免删除键出错,默认设为None super().__init__(*args, **kwargs) def train_step(self, data): # 保留你的train_step实现 pass def get_config(self): config = super().get_config() # 不将pool加入配置(无法序列化) return config @classmethod def from_config(cls, config, custom_objects=None): # 先按基础逻辑重建模型 model = super().from_config(config, custom_objects) # 后续手动赋值pool model.pool = None return model
加载模型时,先加载再手动设置pool:
model = tf.keras.models.load_model("C:\src\models\model", custom_objects={"CustomModel": CustomModel}) # 重新创建多进程Pool并赋值 from multiprocessing import Pool model.pool = Pool(processes=4) # 按实际需求设置进程数
方案2:重建模型结构后加载权重
如果方案1无效,可先重建模型结构,再加载保存的权重:
- 完全复刻训练时的模型结构:
# 重新定义和训练时完全一致的inputs、outputs inputs = ... # 你的输入层定义 outputs = ... # 你的输出层定义 pool = Pool(processes=4) model = CustomModel(inputs=inputs, outputs=outputs, pool=pool) model.compile(optimizer="adam", metrics=["loss"])
- 加载保存的权重:
model.load_weights("C:\src\models\model")
方案3:注册相关指标类解决MeanMetricWrapper报错
若仍有MeanMetricWrapper错误,尝试将相关指标类一并加入custom_objects:
model = tf.keras.models.load_model( "C:\src\models\model", custom_objects={ "CustomModel": CustomModel, "MeanMetricWrapper": tf.keras.metrics.MeanMetricWrapper, "Mean": tf.keras.metrics.Mean } ) # 加载后手动设置pool model.pool = Pool(processes=4)
关键提示
- 多进程
Pool属于不可序列化对象,无法被Keras保存,必须在加载后手动重建并赋值。 - 自定义模型类必须正确实现
get_config和from_config,否则Keras无法正常重建模型实例。
内容的提问来源于stack exchange,提问作者Osterzine

