打印分类报告触发TypeError:真实与预测标签类型不匹配的排查与解决
问题原因与解决方法
报错原因
真实标签y_test的元素是字符串类型(如'0'、'1'),而模型预测结果y_predict的元素是整数类型(如0、1),两者数据类型不匹配。classification_report要求输入的真实标签和预测标签必须为同一数据类型,因此触发该TypeError。
解决方法
只需将两者类型统一即可,以下两种方法任选其一:
方法1:将y_test转换为整数类型
如果y_test是numpy数组:
y_test = y_test.astype(int)
如果y_test是pandas Series:
y_test = y_test.astype('int64')
方法2:将y_predict转换为字符串类型
如果y_predict是numpy数组:
import numpy as np y_predict = y_predict.astype(str)
如果是普通列表:
y_predict = [str(x) for x in y_predict]
调整类型后重新运行print(classification_report(y_test,y_predict))即可正常生成分类报告。
内容的提问来源于stack exchange,提问作者أرْوَى أحْمَد.
相关产品推荐
相关产品推荐

