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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 06:05:00