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

tf.keras加载模型遇MeanMetricWrapper恢复失败求助

Keras模型保存后无法加载的修复方案

问题详情

花费大量时间训练的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_objects arg when calling load_model(). Also, check that the class implements get_config and from_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数据)

加载尝试及对应报错

  1. 指定MeanMetricWrapper加载:
tf.keras.models.load_model("C:\src\models\model", custom_objects={"MeanMetricWrapper": tf.metrics.MeanMetricWrapper})

仍出现相同的MeanMetricWrapper错误。

  1. 指定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无效,可先重建模型结构,再加载保存的权重:

  1. 完全复刻训练时的模型结构:
# 重新定义和训练时完全一致的inputs、outputs
inputs = ...  # 你的输入层定义
outputs = ...  # 你的输出层定义
pool = Pool(processes=4)
model = CustomModel(inputs=inputs, outputs=outputs, pool=pool)
model.compile(optimizer="adam", metrics=["loss"])
  1. 加载保存的权重:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 15:22:48