TensorFlow 2.17模型本地加载报错:MultiHeadAttention参数未识别
英僧伽罗语Transformer模型加载报错解决
问题描述
在Colab环境使用TensorFlow 2.17.0训练的英僧伽罗语Transformer翻译模型可正常运行,但本地同样使用TensorFlow 2.17版本加载该模型时,出现MultiHeadAttention类反序列化错误,提示存在未识别的参数query_shape、key_shape、value_shape。已匹配版本但问题仍存在。
报错栈
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) File ~\AppData\Roaming\Python\Python312\site-packages\keras\src\ops\operation.py:208, in Operation.from_config(cls, config) 207 try: --> 208 return cls(**config) 209 except Exception as e: File ~\AppData\Roaming\Python\Python312\site-packages\keras\src\layers\attention\multi_head_attention.py:115, in MultiHeadAttention.__init__(self, num_heads, key_dim, value_dim, dropout, use_bias, output_shape, attention_axes, kernel_initializer, bias_initializer, kernel_regularizer, bias_regularizer, activity_regularizer, kernel_constraint, bias_constraint, **kwargs) 97 def __init__( 98 self, 99 num_heads, (...) 113 **kwargs, 114 ): --> 115 super().__init__(**kwargs) 116 self.supports_masking = True File ~\AppData\Roaming\Python\Python312\site-packages\keras\src\layers\layer.py:264, in Layer.__init__(self, activity_regularizer, trainable, dtype, autocast, name, **kwargs) 263 if kwargs: --> 264 raise ValueError( 265 "Unrecognized keyword arguments " 266 f"passed to {self.__class__.__name__}: {kwargs}" 267 ) 269 self.built = False ... 213 ) TypeError: Error when deserializing class 'MultiHeadAttention' using config={'name': 'multi_head_attention', 'trainable': True, 'dtype': 'float32', 'num_heads': 4, 'key_dim': 256, 'value_dim': 256, 'dropout': 0.1, 'use_bias': True, 'output_shape': None, 'attention_axes': [1], 'kernel_initializer': {'module': 'keras.initializers', 'class_name': 'GlorotUniform', 'config': {'seed': None}, 'registered_name': None}, 'bias_initializer': {'module': 'keras.initializers', 'class_name': 'Zeros', 'config': {}, 'registered_name': None}, 'kernel_regularizer': None, 'bias_regularizer': None, 'activity_regularizer': None, 'kernel_constraint': None, 'bias_constraint': None, 'query_shape': [None, 55, 256], 'key_shape': [None, 55, 256], 'value_shape': [None, 55, 256]}. Exception encountered: Unrecognized keyword arguments passed to MultiHeadAttention: {'query_shape': [None, 55, 256], 'key_shape': [None, 55, 256], 'value_shape': [None, 55, 256]}
模型代码
from tensorflow.keras.layers import Input, Embedding, MultiHeadAttention, LayerNormalization, Dense, Dropout from tensorflow.keras.models import Model def transformer_encoder(inputs, head_size, num_heads, ff_dim, dropout=0): # Normalization and Attention x = LayerNormalization(epsilon=1e-6)(inputs) x = MultiHeadAttention(key_dim=head_size, num_heads=num_heads, dropout=dropout)(x, x) x = Dropout(dropout)(x) res = x + inputs # Feed Forward Part x = LayerNormalization(epsilon=1e-6)(res) x = Dense(ff_dim, activation="relu")(x) x = Dropout(dropout)(x) x = Dense(inputs.shape[-1])(x) return x + res def build_transformer_model(input_shape, english_vocab_size, sinhala_vocab_size, head_size=256, num_heads=4, ff_dim=512, num_layers=4, dropout=0.1): inputs = Input(shape=input_shape) x = Embedding(english_vocab_size, head_size)(inputs) for _ in range(num_layers): x = transformer_encoder(x, head_size, num_heads, ff_dim, dropout) outputs = Dense(sinhala_vocab_size, activation='softmax')(x) model = Model(inputs, outputs) return model # Build the model input_shape = (max_len,) model = build_transformer_model(input_shape, english_vocab_size, sinhala_vocab_size) # Compile the model model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4), loss='sparse_categorical_crossentropy', metrics=['accuracy']) # Train the model history = model.fit(X_train, y_train, validation_data=(X_val, y_val), epochs=3, batch_size=64) # Save the model model.save('sinhala_english_transformer.h5') import tensorflow as tf from tensorflow.keras.models import load_model from tensorflow.keras.preprocessing.sequence import pad_sequences import pickle # Load the model model = load_model('/content/drive/MyDrive/sinhala_english_transformer_new.h5') # Load the tokenizers with open('/content/drive/MyDrive/english_tokenizer_new.pkl', 'rb') as f: english_tokenizer = pickle.load(f) with open('/content/drive/MyDrive/sinhala_tokenizer_new.pkl', 'rb') as f: sinhala_tokenizer = pickle.load(f) # Set max_len to the value used during training max_len = 55 # Adjust this to the value used during training # Prepare new sentences for prediction (use the tokenizer and padding as before) new_sentences = ["How does the Java Stream API help in processing collections?"] new_sequences = english_tokenizer.texts_to_sequences(new_sentences) new_sequences = pad_sequences(new_sequences, maxlen=max_len, padding='post') # Predict translation predictions = model.predict(new_sequences) translated_sequences = predictions.argmax(axis=-1) translated_sentences = [' '.join(sinhala_tokenizer.index_word[idx] for idx in seq if idx > 0) for seq in translated_sequences] print(translated_sentences)
解决方案
问题根源
Colab环境的TensorFlow 2.17在保存.h5格式模型时,会额外将query_shape、key_shape、value_shape这些运行时参数序列化到模型配置中,但本地环境的MultiHeadAttention层初始化逻辑不接受这些非标准参数,导致反序列化失败。
方法1:自定义MultiHeadAttention反序列化逻辑
通过继承MultiHeadAttention类,重写from_config方法过滤无效参数,再加载模型:
from tensorflow.keras.layers import MultiHeadAttention from tensorflow.keras.models import load_model class CustomMultiHeadAttention(MultiHeadAttention): @classmethod def from_config(cls, config): # 移除不被支持的参数 for param in ['query_shape', 'key_shape', 'value_shape']: config.pop(param, None) return super().from_config(config) # 加载模型时指定自定义对象 model = load_model('sinhala_english_transformer.h5', custom_objects={'MultiHeadAttention': CustomMultiHeadAttention})
方法2:改用SavedModel格式保存与加载
.h5格式属于旧版Keras保存格式,容易出现跨环境兼容问题。改用TensorFlow官方推荐的SavedModel格式,序列化逻辑更稳定:
# 训练后保存为SavedModel格式(替换原有的model.save(.h5)) model.save('sinhala_english_transformer_savedmodel') # 本地加载时直接调用 import tensorflow as tf model = tf.keras.models.load_model('sinhala_english_transformer_savedmodel')
方法3:强制安装完全匹配的TensorFlow版本
虽然表面版本都是2.17,但Colab与本地的Keras子依赖可能存在细微差异,重新安装精确版本:
pip uninstall tensorflow -y pip install tensorflow==2.17.0 --force-reinstall
内容的提问来源于stack exchange,提问作者Kalpani Athukorala
相关产品推荐
相关产品推荐

