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

使用Keras函数式API保存含Masking层模型至.keras文件后加载失败

解决Keras Masking层序列化/加载失败问题

错误原因

问题根源是Masking层的mask_value使用了numpy数组。Keras序列化层配置时,对原生Python类型(如列表)的支持更可靠,而numpy数组作为非原生对象,默认序列化逻辑无法正确解析,导致加载时抛出反序列化错误。


解决方案1:修改模型定义,使用Python列表作为mask_value

直接将MASK_VALUE从numpy数组替换为Python列表,让Keras可以正常序列化Masking层配置:

import tensorflow as tf
from tensorflow.keras.layers import Input, Masking, LSTM, Dropout, Dense, Concatenate, Model

N_FEATURES = 6
# 用Python列表替代numpy数组
MASK_VALUE = [0.0 for _ in range(N_FEATURES)]

def get_clean_model():
    input_layer = Input(shape=(None, N_FEATURES))
    # 使用列表类型的mask_value
    masked_input = Masking(mask_value=MASK_VALUE)(input_layer)
    
    lstm_layer = LSTM(units=N_FEATURES, activation='tanh', return_sequences=True,
                  recurrent_regularizer='l2', kernel_regularizer='l2')(masked_input)
    dropout_layer = Dropout(0.05)(lstm_layer)
    
    dense_layer1 = Dense(N_FEATURES, activation='sigmoid', kernel_regularizer='l2')(dropout_layer)
    dense_layer2 = Dense(N_FEATURES*2, activation='sigmoid', kernel_regularizer='l2')(dense_layer1)
    
    dropout_layer2 = Dropout(0.05)(dense_layer2)
    output_layer = Dense(1, activation='sigmoid')(dropout_layer2)
    
    model = Model(inputs=input_layer, outputs=output_layer)
    return model

# 保存模型
model = get_clean_model()
model.save('model.keras')

# 加载模型
loaded_model = tf.keras.models.load_model('model.keras')

解决方案2:加载已保存的旧模型(无需重新训练)

如果已经用numpy数组作为mask_value保存了模型,可通过自定义对象序列化逻辑加载:

import tensorflow as tf
import numpy as np

# 自定义Masking层的反序列化函数
def custom_masking_from_config(config):
    # 将配置中的mask_value从列表转回numpy数组(按需处理)
    config['mask_value'] = np.asarray(config['mask_value'])
    return tf.keras.layers.Masking(**config)

# 加载模型时注册自定义反序列化逻辑
loaded_model = tf.keras.models.load_model(
    'model.keras',
    custom_objects={'Masking': custom_masking_from_config}
)

也可以先构建模型结构(使用Python列表版MASK_VALUE),再加载权重:

# 先构建模型结构
model = get_clean_model()
# 加载已保存模型的权重
model.load_weights('model.keras')

关键注意点

  • Keras层的配置参数优先使用原生Python类型(如列表、整数、字符串),避免使用numpy数组或其他自定义对象,减少序列化问题。
  • 若必须使用numpy数组,需在custom_objects中指定对应的反序列化逻辑,确保加载时能正确解析参数。

内容的提问来源于stack exchange,提问作者Mark C.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 05:13:27