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

如何实现可定义输入数、带可训练权重的加权和层并解决相关报错

错误原因

你遇到的报错和模型保存问题都是自定义层的写法不规范导致的:

  1. __init__方法中直接对权重变量做除法运算并重新赋值给self.sum_weights,会把原本的tf.Variable类型转换为静态图张量,跨上下文调用时就会触发Graph张量泄漏的报错
  2. call方法中直接修改层的权重变量值,会破坏权重的可训练属性,也会导致图结构异常
  3. 没有实现get_config方法,Keras无法序列化自定义层的参数,保存模型时就会抛出各类序列化错误

修改方案

调整后的代码如下:

class WeightedSum(krs.layers.Layer):
    def __init__(self, n_models=2, **kwargs):
        super(WeightedSum, self).__init__(**kwargs)
        self.n_models = n_models
        # 标准写法声明可训练权重,不在初始化阶段做运算
        w_init = tf.random_uniform_initializer()
        self.sum_weights = self.add_weight(
            shape=(1, self.n_models),
            initializer=w_init,
            trainable=True,
            name="sum_weights"
        )

    def call(self, inputs):
        # 归一化逻辑放在前向传播阶段执行,用临时变量存储结果,不修改原权重
        # 用softmax做归一化比手动除以总和数值更稳定,自动保证权重和为1
        normalized_weights = tf.nn.softmax(self.sum_weights, axis=-1)
        normalized_weights = tf.cast(normalized_weights, dtype=inputs[0].dtype)
        # 矩阵运算替代循环,执行效率更高,适配静态图要求
        inputs_stack = tf.stack(inputs, axis=1)
        output = tf.squeeze(tf.matmul(normalized_weights, inputs_stack), axis=1)
        return output

    # 实现序列化方法,支持模型保存加载
    def get_config(self):
        config = super(WeightedSum, self).get_config()
        config.update({
            "n_models": self.n_models
        })
        return config

加载保存的模型说明

加载包含该自定义层的模型时,需要指定自定义层映射:

model = krs.models.load_model("你的模型路径.h5", custom_objects={"WeightedSum": WeightedSum})

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 13:54:04