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

Keras自定义层保存完整模型时出错,寻求技术解决方案

解决禁用Eager Execution时Keras自定义ConstLayer的保存/加载问题

问题描述

在禁用Eager Execution(tf.compat.v1.disable_eager_execution())的环境下,自定义ConstLayer在模型保存和加载时出现两类错误:

  • 未在get_config中添加config['x'] = self.x:加载模型时报TypeError: __init__()缺少必需的位置参数x
  • 添加上述代码后:保存模型时报NotImplementedError: deepcopy()仅在启用eager execution时可用

解决方案

核心思路是避免在配置中存储Tensor对象,改用可序列化的numpy数组或Python原生数值类型。

方案1:修改get_config存储numpy值

class ConstLayer(tf.keras.layers.Layer):
    def __init__(self, x, **kwargs):
        super(ConstLayer, self).__init__(**kwargs)
        self.x = tf.Variable(x, trainable=False)

    def call(self, input):
        return self.x

    def get_config(self):
        config = super(ConstLayer, self).get_config()
        # 将Tensor转换为numpy数组存储,规避deepcopy Tensor的限制
        config['x'] = self.x.numpy()
        return config

方案2:用Keras后端方法获取值(兼容更多环境)

如果numpy()调用存在兼容性问题,可替换为Keras后端API:

def get_config(self):
    config = super(ConstLayer, self).get_config()
    config['x'] = tf.keras.backend.get_value(self.x)
    return config

方案3:增强__init__的兼容性

确保加载时numpy数组能正确转为Tensor:

class ConstLayer(tf.keras.layers.Layer):
    def __init__(self, x, **kwargs):
        super(ConstLayer, self).__init__(**kwargs)
        # 统一将输入转为Tensor后创建Variable
        self.x = tf.Variable(tf.convert_to_tensor(x), trainable=False)

    def call(self, input):
        return self.x

    def get_config(self):
        config = super(ConstLayer, self).get_config()
        config['x'] = tf.keras.backend.get_value(self.x)
        return config

原理说明

禁用Eager Execution时,Tensor对象无法被序列化(底层依赖的deepcopy操作不支持非Eager模式的Tensor)。通过将Tensor转换为numpy数组存储到配置中,既能满足__init__对参数x的需求,又能避免序列化时的deepcopy错误。加载模型时,Keras会将配置中的numpy数组传入__init__,再转为Tensor创建Variable,与原层的逻辑完全一致。

内容的提问来源于stack exchange,提问作者Mihai.Mehe

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 09:45:28