如何列出测试集中预测错误的样本行?(附现有准确率计算代码)
找出测试集中预测错误的样本行
这很好实现,你只需要通过对比预测结果和真实标签,筛选出不匹配的样本即可。结合你现有的代码,我给你整理了完整的操作步骤:
生成错误标记掩码
先创建一个布尔数组,标记测试集中每一行是否预测错误:# 对比预测值和真实标签,得到错误行的掩码 errors_mask = predicted_mnb_lemmatized != test_data['COMPID']提取错误预测的样本
用这个掩码直接筛选测试集数据,就能得到所有预测错误的样本。如果想同时直观看到预测值,还可以把预测结果添加为新列:# 筛选出所有预测错误的样本(用copy避免修改原数据的警告) wrong_predictions = test_data[errors_mask].copy() # 添加预测结果列,方便对比真实值和预测值 wrong_predictions['PREDICTED_COMPID'] = predicted_mnb_lemmatized[errors_mask]查看错误样本详情
现在你可以直接打印、输出或者交互式查看wrong_predictions,里面包含了所有预测错误的测试样本,以及对应的预测结果,能帮你快速定位模型的问题点。
把这些代码加到你现有的流程里,完整代码大概是这样:
text_mnb_lemmatized = Pipeline([('vect', lemma_count_vect), ('tfidf', TfidfTransformer(sublinear_tf=True, use_idf=False)), ('mnb', MultinomialNB(alpha=0.1, fit_prior=True))]) text_mnb_lemmatized = text_mnb_lemmatized.fit(train_data['CDESCR'], train_data['COMPID']) predicted_mnb_lemmatized = text_mnb_lemmatized.predict(test_data['CDESCR']) # 计算并打印准确率 accuracy = np.mean(predicted_mnb_lemmatized == test_data['COMPID']) print(f"模型准确率: {accuracy:.4f}") # 找出并查看预测错误的样本 errors_mask = predicted_mnb_lemmatized != test_data['COMPID'] wrong_predictions = test_data[errors_mask].copy() wrong_predictions['PREDICTED_COMPID'] = predicted_mnb_lemmatized[errors_mask] print("\n预测错误的样本详情:") print(wrong_predictions)
这样你就能清晰看到哪些测试行预测出错,以及模型给出的错误预测值是什么,方便后续分析模型的不足——比如是不是某些类别样本量太少,或者文本特征提取逻辑需要优化。
内容的提问来源于stack exchange,提问作者juliano.net
相关产品推荐
相关产品推荐

