如何在Keras中定位文本分类验证集误判样本并可视化错误预测
当然可以实现!在Keras里追踪验证集的误判样本并画出你想要的柱状图完全没问题,我给你拆解成几个简单步骤,照着做就行:
第一步:抓取验证集里的误判样本
首先你需要用训练好的模型对验证集做预测,然后和真实标签比对,把预测错的样本挑出来。代码大概是这样的:
# 假设你已经有训练好的model,以及验证集的x_val(输入数据)、y_val(真实标签)、还有原文本texts_val(用来统计长度) y_pred = model.predict(x_val) # 把预测概率转成类别标签(多分类用argmax,二分类也可以用阈值比如0.5) y_pred_classes = np.argmax(y_pred, axis=1) # 如果你的y_val是one-hot编码,就转成整数标签;如果本来就是整数,直接用y_val就行 y_true_classes = np.argmax(y_val, axis=1) # 找出所有误判样本的索引 misclassified_indices = np.where(y_pred_classes != y_true_classes)[0] # 提取对应的原文本、真实标签和预测标签 misclassified_texts = [texts_val[i] for i in misclassified_indices]
这里要注意:如果你的x_val是已经预处理过的token序列,那记得保留原文本texts_val,这样统计出来的句子长度才是真实的;要是你直接用token序列的长度当“句子长度”,那直接取len(x_val[i])就行。
第二步:统计误判样本的长度分布
接下来统计不同长度的误判样本有多少个,用Counter就能轻松搞定:
from collections import Counter # 计算每个误判样本的长度 misclassified_lengths = [len(text) for text in misclassified_texts] # 统计每个长度出现的次数 length_counts = Counter(misclassified_lengths) # 整理成画图需要的格式:X轴是长度(排序后更直观),Y轴是对应次数 x_lengths = sorted(length_counts.keys()) y_counts = [length_counts[length] for length in x_lengths]
第三步:画出柱状图
用Matplotlib就能画出你要的图,代码很简单:
import matplotlib.pyplot as plt plt.figure(figsize=(12, 6)) plt.bar(x_lengths, y_counts, color='#ff6b6b') plt.xlabel('句子长度') plt.ylabel('误判样本数量') plt.title('验证集误判样本的句子长度分布') plt.xticks(rotation=45) # 长度多的话旋转X轴标签,避免重叠 plt.tight_layout() # 自动调整布局 plt.show()
哪怕你训练轮次极少,这个方法也照样能用——它只是客观展示当前模型的误判情况,和模型收敛程度无关,只要有预测结果就能统计。
另外,如果你想在每轮训练结束后都跟踪误判分布,可以自定义一个Keras回调函数,这样能看到训练过程中误判情况的变化:
from keras.callbacks import Callback from collections import Counter import numpy as np class MisclassifiedTracker(Callback): def __init__(self, x_val, y_val, texts_val): super().__init__() self.x_val = x_val self.y_val = y_val self.texts_val = texts_val # 保存原验证文本 def on_epoch_end(self, epoch, logs=None): # 对验证集做预测 y_pred = self.model.predict(self.x_val, verbose=0) y_pred_classes = np.argmax(y_pred, axis=1) y_true_classes = np.argmax(self.y_val, axis=1) # 找出误判样本 misclassified_indices = np.where(y_pred_classes != y_true_classes)[0] # 统计长度分布 misclassified_lengths = [len(self.texts_val[i]) for i in misclassified_indices] length_counts = Counter(misclassified_lengths) # 打印当前轮次的误判分布 print(f"\nEpoch {epoch+1} 误判长度统计: {dict(length_counts)}") # 要是想每轮都保存图,把画图代码放这里就行 # 训练的时候把这个回调加进去 tracker = MisclassifiedTracker(x_val, y_val, texts_val) model.fit(x_train, y_train, validation_data=(x_val, y_val), epochs=5, callbacks=[tracker])
这样每训练一轮,你就能看到当前模型对不同长度句子的误判情况,方便你针对性调整模型(比如给长句子加更多处理逻辑)。
内容的提问来源于stack exchange,提问作者LagSurfer
相关产品推荐
相关产品推荐

