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

BERT情感分析训练AxisError: axis 1越界一维数组报错修复

问题根因

报错AxisError: axis 1 is out of bounds for array of dimension 1的核心触发原因是维度计算逻辑错误,网上通用的改axis=0的方案不生效,是因为代码里同时存在两个维度相关的问题:

  1. 预测结果的argmax轴参数错误:经过np.concatenate拼接后的predictions是形状为(验证集总样本数, 分类类别数)的二维数组,每一行对应一个样本在所有类别上的输出logit。要拿到每个样本的预测类别,需要沿着**类别维度(axis=1)**取最大值索引,写axis=0是沿着样本维度跨行计算,输出形状为(分类类别数,),和标签长度完全不匹配。
  2. 真实标签被错误执行了argmax操作:拼接后的true_vals是形状为(验证集总样本数,)的一维数组,每个元素直接就是对应样本的真实类别ID,本身已经是标签格式,不需要做argmax计算。对一维数组传入axis=1会直接触发轴越界报错;就算传入axis=0,返回的也只是整个标签数组中最大值的索引(单个标量),根本不是逐样本的标签序列,无法计算F1值。
修复方案

直接替换计算F1值的错误代码行即可:

  • 原错误代码:
val_f1 = f1_score_func(np.argmax(predictions, axis=0), np.argmax(true_vals, axis=0))
  • 替换为正确代码:
val_f1 = f1_score_func(np.argmax(predictions, axis=1), true_vals)

可选校验步骤

如果担心维度还有问题,可以在计算F1前加两行打印语句确认数组形状:

# 校验维度
print(f"拼接后预测结果形状: {predictions.shape}")
print(f"拼接后真实标签形状: {true_vals.shape}")
print(f"argmax后预测标签形状: {np.argmax(predictions, axis=1).shape}")

正常输出会符合如下格式(以1000条验证样本、3分类任务为例):

拼接后预测结果形状: (1000, 3)
拼接后真实标签形状: (1000,)
argmax后预测标签形状: (1000,)

预测标签和真实标签形状完全对齐,即可正常计算加权F1值。

其他非阻断问题修正

训练循环里的训练loss打印逻辑存在计算误差,不影响运行但显示数值不准:

# 原代码 len(batch)返回的是batch内张量个数(固定为3,对应input_ids、attention_mask、labels),不是batch size
progress_bar.set_postfix({'training_loss': '{:.3f}'.format(loss.item()/len(batch))})

# 修正为除以当前batch的样本量
progress_bar.set_postfix({'training_loss': '{:.3f}'.format(loss.item()/batch[0].size(0))})

内容的提问来源于stack exchange,提问作者Vincent Van

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 16:15:12