在LSTM模型上使用LIME时出现错误,求排查(附代码与错误截图)
问题分析与解决:LIME适配LSTM模型时的输入格式错误
错误原因
核心问题是LIME与LSTM的输入格式不兼容:
LimeTextExplainer默认针对原始文本字符串工作,会对文本进行单词级扰动生成解释样本;但你的ls_X_test[idx]大概率是经过Tokenizer预处理后的数字序列(而非原始文本),LIME无法识别这种格式完成扰动操作。- LSTM模型的
predict_proba方法要求输入是批量的token化数字序列(形状通常为(batch_size, max_seq_len)),但LIME生成的扰动样本是原始文本字符串,直接传入模型会引发类型/维度不匹配错误。
解决步骤
1. 准备原始文本输入
确保传入LIME的是未经过预处理的原始文本:
- 若你有单独的原始测试文本列表
raw_X_test,直接使用该列表; - 若只有token化后的
ls_X_test,可通过Tokenizer还原文本(注意处理OOV标记):raw_X_test = tokenizer.sequences_to_texts(ls_X_test)
2. 封装适配LSTM的预测函数
编写包装函数,将LIME生成的原始文本先转为模型可接受的token化序列:
# 假设已定义tokenizer、LSTM模型c,以及训练时设置的MAX_SEQ_LENGTH def lstm_predict_proba(texts): # 将文本列表转为token化序列 tokenized = tokenizer.texts_to_sequences(texts) # 填充序列到模型要求的固定长度 padded = pad_sequences(tokenized, maxlen=MAX_SEQ_LENGTH) # 调用模型输出概率 return c.predict_proba(padded)
3. 修改LIME调用代码
使用封装后的预测函数,传入原始文本完成解释:
# 初始化LIME文本解释器 lime_explainer = LimeTextExplainer(class_names=class_names) # 生成解释结果 LIME_exp = lime_explainer.explain_instance(raw_X_test[idx], lstm_predict_proba, num_features=5) # 打印结果(统一使用适配后的预测函数) print('Document id: %d' % idx) print('Text: ', raw_X_test[idx]) print('Probability d =', lstm_predict_proba([raw_X_test[idx]]).round(3)[0,1]) print('True class: %s' % class_names.get(list(y_test)[idx]))
额外注意事项
- 确保
MAX_SEQ_LENGTH与训练LSTM时的参数完全一致; - 若处理中文文本,需通过
LimeTextExplainer的split_expression参数设置自定义分词逻辑(比如jieba分词),保证和Tokenizer的分词规则匹配。
内容的提问来源于stack exchange,提问作者Healer_524
相关产品推荐
相关产品推荐

