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

PyTorch中如何获取测试集分类错误图像的文件名?

问题解决步骤与代码调整

核心问题分析

  1. 错误根源:你尝试用torch.cat()处理字符串列表(文件名),但该函数仅支持张量操作;同时用批次索引i去索引数据集的DataFrame,会导致仅获取每个批次的第一个文件名,且在shuffle模式下文件名与样本不匹配。
  2. 关键改进方向:让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 08:30:50