PyTorch计算模型准确率出现TypeError: 'bool' object is not iterable错误
报错原因及修复说明
核心报错根源
- 首先是
torch.max返回值的误用:调用torch.max(outputs.data, 1)时,返回的是二元组(对应维度最大值张量, 最大值对应索引张量),你需要的预测类别是第二个返回的索引张量。之前直接把整个二元组赋值给predicted后做predicted == label运算,只会得到单个布尔值,而非逐样本对比的布尔张量。 - 其次是运算顺序错误:此前的写法
float((predicted == label)).sum()先把得到的单个布尔值强转为float类型(True转1.0、False转0.0),对0维的数值变量调用sum()或者直接传入sum()迭代,自然触发“bool对象不可迭代”的报错。
修复逻辑
你调整后的代码是标准正确写法:
- 先拆分
torch.max的返回值,仅保留预测类别索引pred_indices - 用
pred_indices == label得到逐样本对比的布尔张量,对该张量求和即可得到当前batch预测正确的样本总数 - 最后调用
.item()是将PyTorch张量类型的结果转为Python原生数值,避免后续累加出现类型兼容问题
额外优化提示
你自定义的Dataset类存在冗余计算:__init__方法中已经把所有y_train转换为对应索引存入self.y_train张量,__getitem__方法不需要再重复调用emotion_list.index(y_train[ind])计算索引,直接返回self.y_train[ind]即可,可减少不必要的性能开销。
内容的提问来源于stack exchange,提问作者John Harrington
相关产品推荐
相关产品推荐

