PyTorch点云分类测试样本数不匹配报错:[908,9080]
问题分析与解决方案
错误根源
报错ValueError: Found input variables with inconsistent numbers of samples: [908, 9080]的核心原因是**y_true和y_pred的数据结构完全不匹配**:
y_true存储的是908个样本的10维one-hot标签(每个样本对应一个长度为10的列表),逻辑上是908个样本;y_pred通过reshape(-1)被强制展平成一维数组(9080个独立元素),被判定为9080个样本;accuracy_score会根据列表顶层元素数量判定样本数,因此出现样本数不匹配的错误。
针对性修复
点云分类任务绝大多数是单标签多分类(每个样本仅属于一个类别),推荐以下修复方案:
修复后的代码(单标签多分类场景)
import torch from sklearn.metrics import accuracy_score def test(model, test_loader): model.eval() y_true = [] y_pred = [] with torch.no_grad(): for data in test_loader: inputs, target = data['pointcloud'].to(device).float(), data['category'].to(device) # 单标签分类无需将标签转为one-hot,直接用原始类别索引(0-9) output_orig = model(inputs) # 单标签多分类用softmax输出类别概率,argmax获取预测类别索引 output = torch.softmax(output_orig, dim=1) pred = torch.argmax(output, dim=1).detach().cpu().numpy() target = target.cpu().numpy() # 直接扩展类别索引列表,两者长度均为908 y_true.extend(target.tolist()) y_pred.extend(pred.tolist()) return accuracy_score(y_true, y_pred)
如果是多标签分类场景(每个样本可属于多个类别)
若你的任务确实是多标签分类,只需移除reshape(-1),让y_pred保持与y_true一致的二维结构即可:
import torch from sklearn.metrics import accuracy_score def test(model, test_loader): model.eval() y_true = [] y_pred = [] with torch.no_grad(): for data in test_loader: inputs, target = data['pointcloud'].to(device).float(), data['category'].to(device) target = torch.nn.functional.one_hot(target, num_classes=10).cpu().float() output_orig = model(inputs) output = torch.sigmoid(output_orig) pred = (output.detach().cpu().numpy() > 0.5) * 1 # 移除reshape(-1),保持每个样本对应10维预测结果 y_true.extend(target.tolist()) y_pred.extend(pred.tolist()) # 多标签场景下,accuracy_score会计算每个样本的标签匹配准确率后取平均 return accuracy_score(y_true, y_pred)
关键注意事项
- 单标签多分类任务中,模型最后一层无需使用sigmoid,应配合softmax+交叉熵损失(
CrossEntropyLoss)训练; - 多标签分类任务才需要使用sigmoid+二元交叉熵损失(
BCELoss)训练,需确保训练与测试逻辑一致。
内容的提问来源于stack exchange,提问作者MrFoxs
相关产品推荐
相关产品推荐

