PyTorch中如何获取测试集分类错误图像的文件名?
问题解决步骤与代码调整
核心问题分析
- 错误根源:你尝试用
torch.cat()处理字符串列表(文件名),但该函数仅支持张量操作;同时用批次索引i去索引数据集的DataFrame,会导致仅获取每个批次的第一个文件名,且在shuffle模式下文件名与样本不匹配。 - 关键改进方向:让Dataset直接返回每个样本的文件名,在迭代DataLoader时同步获取对应批次的文件名,避免通过索引硬匹配。
步骤1:修复Dataset定义与初始化
调整CustomDataset类
修改__getitem__方法,返回文件名;同时优化标签格式为标量张量,便于后续处理:
class CustomDataset(Dataset): def __init__(self, img_path, csv_file, transforms): self.imgs_path = img_path self.csv_train_file = csv_file self.data_df = pd.read_csv(self.csv_train_file) self.transforms = transforms self.data = [] self.class_map = {"ProbableAD" : 0, "Control": 1} self.img_dim = (256, 256) for ind in self.data_df.index: img_path = self.data_df['spectrogramSegFilename'][ind] class_name = self.data_df['dx'][ind] self.data.append([img_path, class_name]) def __len__(self): return len(self.data) def __getitem__(self, idx): img_path, class_name = self.data[idx] img = cv2.imread(img_path) img = cv2.resize(img, self.img_dim) class_id = self.class_map[class_name] # 图像预处理 img_tensor = torch.from_numpy(img).permute(2, 0, 1) data = self.transforms(img_tensor) # 返回文件名+标量标签(避免后续维度冗余) return data, torch.tensor(class_id, dtype=torch.long), img_path
修复主函数中的数据集初始化
删除重复定义的无transforms参数的数据集:
if __name__ == "__main__": transformations = transforms.Compose([ transforms.ToPILImage(), transforms.Resize(256), transforms.CenterCrop(256), transforms.ToTensor(), transforms.Normalize((0.49966475, 0.1840554, 0.34930056), (0.35317238, 0.17343724, 0.1894943)) ]) # 仅保留正确的数据集定义 train_dataset = CustomDataset("/spectrogram_images/spectrogram_train/", "train_features_segmented.csv", transformations) test_dataset = CustomDataset("/spectrogram_images/spectrogram_test/", "test_features_segmented.csv", transformations) train_data_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_data_loader = DataLoader(test_dataset, batch_size=64, shuffle=True)
步骤2:重写get_predictions函数
该函数将同步收集文件名、真实标签、预测标签与概率:
def get_predictions(model, iterator, device): model.eval() filenames = [] true_labels = [] probs = [] with torch.no_grad(): # 迭代时获取每个批次的图像、标签、文件名 for x, y, batch_filenames in iterator: x = x.to(device) y_pred = model(x) y_prob = F.softmax(y_pred, dim=-1) # 收集当前批次的文件名、标签、概率 filenames.extend(batch_filenames) true_labels.append(y.cpu()) probs.append(y_prob.cpu()) # 拼接张量并计算预测标签 true_labels = torch.cat(true_labels, dim=0) probs = torch.cat(probs, dim=0) pred_labels = torch.argmax(probs, dim=1) return filenames, true_labels, pred_labels, probs
步骤3:生成结果DataFrame并筛选错误分类
# 假设model和device已定义 filenames, true_labels, pred_labels, probs = get_predictions(model, test_data_loader, device) # 转换为numpy数组便于DataFrame处理 true_labels = true_labels.numpy() pred_labels = pred_labels.numpy() prob_ad = probs[:, 0].numpy() prob_control = probs[:, 1].numpy() # 创建结果DataFrame results_df = pd.DataFrame({ 'filename': filenames, 'true_label': true_labels, 'pred_label': pred_labels, 'prob_ProbableAD': prob_ad, 'prob_Control': prob_control }) # 映射标签编号到类别名称(可选) results_df['true_label'] = results_df['true_label'].map({0: 'ProbableAD', 1: 'Control'}) results_df['pred_label'] = results_df['pred_label'].map({0: 'ProbableAD', 1: 'Control'}) # 筛选分类错误的样本 misclassified_df = results_df[results_df['true_label'] != results_df['pred_label']] # 保存结果(可选) misclassified_df.to_csv("misclassified_spectrograms.csv", index=False)
内容的提问来源于stack exchange,提问作者csStudent2102
相关产品推荐
相关产品推荐

