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

PyTorch自定义数据集加载报错:文件未找到与XML解析失败

问题解决:PyTorch自定义Dataset加载数据集的路径与解析错误

错误原因分析

  1. 第一个FileNotFoundError:
    遍历root_dir时,os.listdir返回的是带后缀的完整文件名(比如00000742_jpg.rf.f14a935bb2e280f6b6e3474673251b7f.jpg),但你额外追加了.jpg和.xml,导致生成路径变成xxx.jpg.jpg和xxx.jpg.xml,自然找不到对应文件。

  2. 第二个ParseError:
    改成label_path = filename后,程序会尝试用XML解析器读取图片文件(遍历到的filename可能是.jpg格式),图片的二进制内容不是合法XML结构,因此触发解析错误。

修复方案

修改__init__方法的文件遍历逻辑,只处理图片文件,提取不带后缀的文件名匹配对应XML标签,同时增加文件存在性检查提升鲁棒性:

修改后的完整代码

import os
import xml.etree.ElementTree as ET
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
import cv2

class CustomDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        self.root_dir = root_dir
        self.transform = transform
        self.images = []
        self.labels = []

        # 仅处理.jpg格式的图片文件
        for filename in os.listdir(root_dir):
            if not filename.endswith('.jpg'):
                continue  # 跳过非图片文件
            
            # 提取不带后缀的文件名,用于匹配XML标签
            file_basename = os.path.splitext(filename)[0]
            image_path = os.path.join(root_dir, filename)
            label_path = os.path.join(root_dir, f"{file_basename}.xml")

            # 检查标签文件是否存在,避免报错
            if not os.path.exists(label_path):
                print(f"警告:未找到{image_path}对应的标签文件{label_path},跳过该样本")
                continue

            # 读取图片并检查是否读取成功
            image = cv2.imread(image_path)
            if image is None:
                print(f"警告:无法读取图片{image_path},跳过该样本")
                continue

            # 解析XML标签
            tree = ET.parse(label_path)
            root = tree.getroot()
            objects = root.findall('object')
            labels = []
            for obj in objects:
                bndbox = obj.find('bndbox')
                xmin = int(bndbox.find('xmin').text)
                ymin = int(bndbox.find('ymin').text)
                xmax = int(bndbox.find('xmax').text)
                ymax = int(bndbox.find('ymax').text)
                labels.append([xmin, ymin, xmax, ymax])
            
            self.images.append(image)
            self.labels.append(labels)

    def __len__(self):
        return len(self.images)

    def __getitem__(self, idx):
        image = self.images[idx]
        label = self.labels[idx]

        if self.transform:
            image = self.transform(image)
        
        return image, label

关键改动说明

  • 筛选图片文件:通过filename.endswith('.jpg')跳过XML或其他无关文件,避免无效遍历。
  • 提取文件基名:用os.path.splitext(filename)[0]获取不带后缀的文件名,确保能正确匹配对应的XML标签文件。
  • 增加存在性检查:检查图片和标签文件是否存在,遇到损坏或缺失文件时打印警告并跳过,防止程序直接崩溃。

内容的提问来源于stack exchange,提问作者Jeremy Ifrah

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 14:02:09