You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python生成混淆矩阵时AxisError axis 1越界错误该如何解决?

报错原因及解决方法

核心错误根因

报错触发的原因是你调用np.argmax(y_val_class, axis=1)时,输入的y_val_class本身已经是1维数组。numpy数组的维度从0开始计数,1维数组仅存在axis=0,指定axis=1属于维度越界。

axis=1的逻辑说明

你原本使用np.argmax(xxx, axis=1)的目的是将shape为(样本数, 类别数)的2维独热编码/概率输出,转换为shape为(样本数,)的1维类别序号:axis=1代表沿着数组的第二个维度(也就是类别维度)计算最大值对应的索引,每个样本生成1个分类结果。
现在y_val_class是1维,说明它已经是预处理完成的分类标签,不需要再做argmax转换。

修复方案

  • 直接删除报错行y_val_class = np.argmax(y_val_class, axis=1),直接用原始的y_val_class输入混淆矩阵即可。
  • 执行前可以先打印数组维度确认格式是否符合预期:
# 维度校验代码
print(y_val_class.shape) # 正常输出应为 (样本数,) 即1维
print(predictions.shape) # 正常输出应为 (样本数, 8) 即2维,保留后续的argmax操作即可
  • 剩余的y_pred_class = np.argmax(predictions, axis=1)逻辑正确,无需修改,模型输出的预测结果是每个样本对应8类的概率,符合axis=1的使用条件。

内容的提问来源于stack exchange,提问作者Safkat Azad ARNAB

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.06 21:39:02