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

无法保存含双编码器的Keras子类模型,询问加载方案可行性

关于DualEncoder模型重建的问题

你给出的写法不可行,核心问题在于__init__方法里硬编码加载本地路径的模型,完全忽略了方法参数中传入的text_encoder和image_encoder,既不灵活,也不符合类的设计逻辑——如果后续需要更换模型路径、或者传入自定义的编码器实例,这种写法会直接失效。

下面是两种正确的重建思路:

方案一:外部加载子模型后传入(推荐)

保持DualEncoder的__init__参数设计,先在类外部加载好两个子模型,再传入初始化:

class DualEncoder(keras.Model):
    def __init__(self, text_encoder, image_encoder, temperature=1.0, **kwargs):
        super(DualEncoder, self).__init__(**kwargs)
        self.text_encoder = text_encoder
        self.image_encoder = image_encoder
        self.temperature = temperature
        self.loss_tracker = keras.metrics.Mean(name="loss")

# 重建步骤
loaded_text_encoder = tf.keras.models.load_model('text_encoder')
loaded_image_encoder = tf.keras.models.load_model('image_encoder')
dual_encoder = DualEncoder(loaded_text_encoder, loaded_image_encoder)

方案二:类内部支持路径加载(灵活扩展)

如果希望DualEncoder自身支持从路径加载子模型,可以修改__init__逻辑,同时支持传入模型实例或路径:

class DualEncoder(keras.Model):
    def __init__(self, text_encoder=None, image_encoder=None, text_encoder_path=None, image_encoder_path=None, temperature=1.0, **kwargs):
        super(DualEncoder, self).__init__(**kwargs)
        # 优先使用传入的模型实例,无实例则从路径加载
        if text_encoder is not None:
            self.text_encoder = text_encoder
        elif text_encoder_path is not None:
            self.text_encoder = tf.keras.models.load_model(text_encoder_path)
        else:
            raise ValueError("必须提供text_encoder实例或text_encoder_path路径")
        
        if image_encoder is not None:
            self.image_encoder = image_encoder
        elif image_encoder_path is not None:
            self.image_encoder = tf.keras.models.load_model(image_encoder_path)
        else:
            raise ValueError("必须提供image_encoder实例或image_encoder_path路径")
        
        self.temperature = temperature
        self.loss_tracker = keras.metrics.Mean(name="loss")

# 通过路径重建
dual_encoder = DualEncoder(text_encoder_path='text_encoder', image_encoder_path='image_encoder')

额外注意事项

  • 如果你的text_encoder或image_encoder是自定义子类模型,加载时需要通过custom_objects参数指定自定义类,示例:
    tf.keras.models.load_model('text_encoder', custom_objects={'TextEncoder': TextEncoder})
    
  • 若之前是用save_weights保存的子模型,需先创建对应子模型的实例,再调用load_weights方法加载权重,而不是使用load_model。

内容的提问来源于stack exchange,提问作者albert

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 18:45:21