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
相关产品推荐
相关产品推荐

