预训练Keras-LSTM模型加载后如何实现新样本分类?
解决加载预训练LSTM模型后预测异常的问题
问题根源分析
- Tokenizer不匹配:训练时的Tokenizer是基于完整训练数据集拟合出的词汇表,你加载模型后重新初始化Tokenizer并仅拟合了测试句子,导致同一单词在训练和预测阶段的编码完全不同,模型无法识别有效输入。
- 缺失分类映射关系:未保存训练时生成的
class_to_index和index_to_class字典,无法将模型输出的概率数组转换为对应的分类标签。 - 预处理参数不一致:预测时的
pad_sequences未指定truncating='post'和padding='post',与训练阶段的预处理逻辑不符,影响输入有效性。
解决方案步骤
步骤1:训练阶段保存必要辅助组件
在训练代码末尾添加序列化代码,保存Tokenizer、分类映射和maxlen参数:
import pickle # 保存训练好的Tokenizer with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/tokenizer.pkl', 'wb') as f: pickle.dump(tokenizer, f) # 保存分类映射字典 with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/class_mappings.pkl', 'wb') as f: pickle.dump({ 'class_to_index': class_to_index, 'index_to_class': index_to_class }, f) # 保存maxlen参数 with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/maxlen.pkl', 'wb') as f: pickle.dump(maxlen, f)
步骤2:加载模型及辅助组件
在预测代码中,先加载模型,再读取保存的Tokenizer、分类映射和maxlen:
import tensorflow as tf import numpy as np import pickle from keras_preprocessing.sequence import pad_sequences # 加载预训练模型 model = tf.keras.models.load_model('/content/drive/MyDrive/Proyect/BehaviorClassifier/twitterBehaviorClassifier.h5') # 加载训练时的Tokenizer with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/tokenizer.pkl', 'rb') as f: tokenizer = pickle.load(f) # 加载分类映射关系 with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/class_mappings.pkl', 'rb') as f: mappings = pickle.load(f) class_to_index = mappings['class_to_index'] index_to_class = mappings['index_to_class'] # 加载maxlen参数 with open('/content/drive/MyDrive/Proyect/BehaviorClassifier/maxlen.pkl', 'rb') as f: maxlen = pickle.load(f)
步骤3:正确预处理新文本并预测
使用训练时的Tokenizer处理新句子,保持预处理逻辑与训练阶段一致:
new_sentence = ["I am very happy"] # 生成符合训练标准的序列 seq = tokenizer.texts_to_sequences(new_sentence) # 按照训练时的参数做padding和截断 padded = pad_sequences(seq, truncating='post', padding='post', maxlen=maxlen) # 获取预测概率 pred_probs = model.predict(padded)[0] # 转换为对应的分类标签 pred_class = index_to_class[np.argmax(pred_probs)] print(f"输入句子: {new_sentence[0]}") print(f"预测分类: {pred_class}") print(f"各分类概率: {pred_probs}")
原代码异常原因说明
你之前重新初始化的Tokenizer仅认识测试句子中的几个单词,生成的序列编码与训练阶段完全不匹配,模型无法理解输入语义,因此输出的是无意义的概率分布。
内容的提问来源于stack exchange,提问作者sigma5563
相关产品推荐
相关产品推荐

