Python2转Python3矩阵报错求助及混淆矩阵代码适配需求
Python2转Python3:混淆矩阵打印代码的适配与错误修复
问题分析
从你给出的代码片段来看,主要存在两类问题:Python2到Python3的语法兼容问题,以及混淆矩阵打印逻辑的未完成部分。我会一步步帮你梳理并修复:
适配与修复后的完整代码
predicted_labels = clf.predict(X_test) def print_cm(cm, labels, hide_zeroes=False, hide_diagonal=False, hide_threshold=None): """Pretty print for confusion matrices (Python3 compatible)""" # 计算每列的宽度:取标签最大长度和数值默认长度(5)的最大值 columnwidth = max([len(str(label)) for label in labels] + [5]) empty_cell = " " * columnwidth # 打印表头(列标签) print(" " + empty_cell, end="") for label in labels: print("%{0}s".format(columnwidth) % label, end="") print() # 换行 # 打印每一行(行标签+对应矩阵值) for i, label1 in enumerate(labels): # 打印行标签 print(" %{0}s".format(columnwidth) % label1, end="") # 遍历当前行的每个数值 for j, label2 in enumerate(labels): cell = "%{0}d".format(columnwidth) % cm[i, j] # 根据规则隐藏指定内容 if hide_zeroes and cm[i, j] == 0: cell = empty_cell if hide_diagonal and i == j: cell = empty_cell if hide_threshold is not None and cm[i, j] < hide_threshold: cell = empty_cell print(cell, end="") print() # 每行结束后换行
关键修改点说明
Python3
print函数适配:- 原代码中
print ("...",)在Python2是实现不换行,但Python3会把逗号后的内容当作元组打印,改用end=""来控制不换行,符合Python3的语法规范。 - 原代码中无括号的
print是Python2语法,Python3必须写成print()。
- 原代码中
逻辑补全:
- 补全了缺失的列遍历循环
for j, label2 in enumerate(labels),实现了混淆矩阵每个单元格数值的打印。 - 完善了
hide_zeroes、hide_diagonal、hide_threshold三个参数的逻辑,确保符合函数注释的预期。
- 补全了缺失的列遍历循环
鲁棒性优化:
- 计算列宽时用
len(str(label))替代len(x),避免标签不是字符串类型时出错。 - 明确了矩阵取值
cm[i, j],确保和混淆矩阵的维度对应(比如sklearn的confusion_matrix返回的二维数组格式)。
- 计算列宽时用
使用示例
假设你已经通过sklearn生成了混淆矩阵:
from sklearn.metrics import confusion_matrix # 假设y_test是真实标签,predicted_labels是预测标签 cm = confusion_matrix(y_test, predicted_labels) # 调用打印函数 print_cm(cm, labels=["classA", "classB", "classC"], hide_zeroes=True)
内容的提问来源于stack exchange,提问作者Alexandra Espinosa
相关产品推荐
相关产品推荐

