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

TensorFlow 2.2.0下BiLSTM+Attention模型架构无法保存问题求助

解决自定义Attention层无法序列化的问题

这个错误的根源很清晰:当自定义Keras层且在__init__中定义了额外参数时,必须重写get_config方法。Keras在序列化(保存)模型架构时,需要通过这个方法获取层的初始化参数,否则调用model.to_json()就会触发NotImplementedError。

具体解决方案:给AttentionWithContext层添加get_config方法

你需要在自定义层类中实现get_config,把__init__里的所有自定义参数都包含进去,同时调用父类的get_config来保留基础配置。修改后的完整代码如下:

from keras import backend as K
def dot_product(x, kernel):
    if K.backend() == 'tensorflow':
        return K.squeeze(K.dot(x, K.expand_dims(kernel)), axis=-1)
    else:
        return K.dot(x, kernel)

class AttentionWithContext(tf.keras.layers.Layer):
    """
    # Input shape
    3D tensor with shape: `(samples, steps, features)`.
    # Output shape
    2D tensor with shape: `(samples, features)`.
    """
    def __init__(self, W_regularizer=None, u_regularizer=None, b_regularizer=None,
                 W_constraint=None, u_constraint=None, b_constraint=None,
                 bias=True, **kwargs):
        self.supports_masking = True
        self.init = tf.keras.initializers.get('glorot_uniform')
        self.W_regularizer = tf.keras.regularizers.get(W_regularizer)
        self.u_regularizer = tf.keras.regularizers.get(u_regularizer)
        self.b_regularizer = tf.keras.regularizers.get(b_regularizer)
        self.W_constraint = tf.keras.constraints.get(W_constraint)
        self.u_constraint = tf.keras.constraints.get(u_constraint)
        self.b_constraint = tf.keras.constraints.get(b_constraint)
        self.bias = bias
        super(AttentionWithContext, self).__init__(**kwargs)

    def build(self, input_shape):
        assert len(input_shape) == 3
        self.W = self.add_weight(shape=(input_shape[-1], input_shape[-1],),
                                 initializer=self.init,
                                 name='{}_W'.format(self.name),
                                 regularizer=self.W_regularizer,
                                 constraint=self.W_constraint)
        if self.bias:
            self.b = self.add_weight(shape=(input_shape[-1],),
                                     initializer='zero',
                                     name='{}_b'.format(self.name),
                                     regularizer=self.b_regularizer,
                                     constraint=self.b_constraint)
        self.u = self.add_weight(shape=(input_shape[-1],),
                                 initializer=self.init,
                                 name='{}_u'.format(self.name),
                                 regularizer=self.u_regularizer,
                                 constraint=self.u_constraint)
        super(AttentionWithContext, self).build(input_shape)

    def compute_mask(self, input, input_mask=None):
        # do not pass the mask to the next layers
        return None

    def call(self, x, mask=None):
        uit = dot_product(x, self.W)
        if self.bias:
            uit += self.b
        uit = K.tanh(uit)
        ait = dot_product(uit, self.u)

        a = K.exp(ait)

        # apply mask after the exp. will be re-normalized next
        if mask is not None:
            # Cast the mask to floatX to avoid float64 upcasting in theano
            a *= K.cast(mask, K.floatx())

        # in some cases especially in the early stages of training the sum may be almost zero
        # and this results in NaN's. A workaround is to add a very small positive number ε to the sum.
        # a /= K.cast(K.sum(a, axis=1, keepdims=True), K.floatx())
        a /= K.cast(K.sum(a, axis=1, keepdims=True) + K.epsilon(), K.floatx())
        a = K.expand_dims(a)
        weighted_input = x * a
        return K.sum(weighted_input, axis=1)

    def compute_output_shape(self, input_shape):
        return input_shape[0], input_shape[-1]
    
    # 新增的get_config方法,用于序列化层参数
    def get_config(self):
        config = super(AttentionWithContext, self).get_config()
        config.update({
            'W_regularizer': tf.keras.regularizers.serialize(self.W_regularizer),
            'u_regularizer': tf.keras.regularizers.serialize(self.u_regularizer),
            'b_regularizer': tf.keras.regularizers.serialize(self.b_regularizer),
            'W_constraint': tf.keras.constraints.serialize(self.W_constraint),
            'u_constraint': tf.keras.constraints.serialize(self.u_constraint),
            'b_constraint': tf.keras.constraints.serialize(self.b_constraint),
            'bias': self.bias
        })
        return config

为什么要这么做?

Keras的模型序列化(比如to_json)需要保存每个层的完整配置信息,这样后续加载模型时才能精准重建每个层。默认的get_config方法只会保留父类的基础参数,不会包含你自定义的正则化、约束、bias等参数,所以必须手动实现这个方法,把所有必要参数序列化进去。

修改后的验证

完成修改后,再调用model.to_json()就能正常生成模型的JSON架构文件了。后续如果需要加载模型,只要确保AttentionWithContext类已经定义(或者注册到Keras自定义层中),就可以用tf.keras.models.model_from_json成功重建模型。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 08:07:37