PyTorch多数据集合并问题:自定义韩文数据集路径错误致文件找不到
问题
训练一款可识别手写韩文、英文及数字的AI模型时,需要合并自定义韩文数据集、MJSynth和SynthText三个数据集。但合并后,train_set会默认使用MJSynth的路径,导致自定义数据集中的图片(如긴장_1227682.jpg)被错误地到MJSynth目录下查找,引发FileNotFoundError。
代码
custom_train_set = RecognitionDataset( parts[0].joinpath("images"), parts[0].joinpath("labels.json"), img_transforms=Compose( [ T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), # Augmentations T.RandomApply(T.ColorInversion(), 0.1), ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.02), ] ), ) if len(parts) > 1: for subfolder in parts[1:]: custom_train_set.merge_dataset( RecognitionDataset(subfolder.joinpath("images"), subfolder.joinpath("labels.json")) ) train_set = MJSynth( train=True, img_folder='/media/cvpr/CM_22/mjsynth/mnt/ramdisk/max/90kDICT32px', label_path='/media/cvpr/CM_22/mjsynth/mnt/ramdisk/max/90kDICT32px/imlist.txt', img_transforms=T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), ) _train_set = SynthText( train=True, recognition_task=True, download=True, # NOTE: download can take really long depending on your bandwidth img_transforms=T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), ) train_set.data.extend([(np_img, target) for np_img, target in _train_set.data]) train_set.data.extend([(np_img, target) for np_img, target in custom_train_set.data])
报错信息
Traceback (most recent call last): File "/media/cvpr/CM_22/doctr/references/recognition/train_pytorch.py", line 485, in <module> main(args) File "/media/cvpr/CM_22/doctr/references/recognition/train_pytorch.py", line 396, in main fit_one_epoch(model, train_loader, batch_transforms, optimizer, scheduler, mb, amp=args.amp) File "/media/cvpr/CM_22/doctr/references/recognition/train_pytorch.py", line 118, in fit_one_epoch for images, targets in progress_bar(train_loader, parent=mb): File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/fastprogress/fastprogress.py", line 50, in __iter__ raise e File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/fastprogress/fastprogress.py", line 41, in __iter__ for i,o in enumerate(self.gen): File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 628, in __next__ data = self._next_data() File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 1333, in _next_data return self._process_data(data) File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 1359, in _process_data data.reraise() File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/_utils.py", line 543, in reraise raise exception FileNotFoundError: Caught FileNotFoundError in DataLoader worker process 0. Original Traceback (most recent call last): File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/_utils/worker.py", line 302, in _worker_loop data = fetcher.fetch(index) File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/_utils/fetch.py", line 58, in fetch data = [self.dataset[idx] for idx in possibly_batched_index] File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/torch/utils/data/_utils/fetch.py", line 58, in <listcomp> data = [self.dataset[idx] for idx in possibly_batched_index] File "/media/cvpr/CM_22/doctr/doctr/datasets/datasets/base.py", line 48, in __getitem__ img, target = self._read_sample(index) File "/media/cvpr/CM_22/doctr/doctr/datasets/datasets/pytorch.py", line 37, in _read_sample else read_img_as_tensor(os.path.join(self.root, img_name), dtype=torch.float32) File "/media/cvpr/CM_22/doctr/doctr/io/image/pytorch.py", line 52, in read_img_as_tensor pil_img = Image.open(img_path, mode="r").convert("RGB") File "/home/cvpr/anaconda3/envs/pytesseract/lib/python3.9/site-packages/PIL/Image.py", line 2912, in open fp = builtins.open(filename, "rb") FileNotFoundError: [Errno 2] No such file or directory: '/media/cvpr/CM_22/mjsynth/mnt/ramdisk/max/90kDICT32px/긴장_1227682.jpg'
解决方案
问题核心是你直接把其他数据集的data元素追加到MJSynth实例中,但MJSynth的读取逻辑会固定拼接自身的root路径(MJSynth图片目录),而自定义数据集和SynthText的图片并不在该目录下,因此报错。
方法1:用ConcatDataset合并数据集(推荐)
PyTorch的ConcatDataset可以直接合并多个独立的Dataset实例,每个实例会用自己的路径逻辑读取图片,不会互相干扰:
from torch.utils.data import ConcatDataset # 保留原有的三个数据集初始化代码 custom_train_set = RecognitionDataset( parts[0].joinpath("images"), parts[0].joinpath("labels.json"), img_transforms=Compose( [ T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), T.RandomApply(T.ColorInversion(), 0.1), ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.02), ] ), ) if len(parts) > 1: for subfolder in parts[1:]: custom_train_set.merge_dataset( RecognitionDataset(subfolder.joinpath("images"), subfolder.joinpath("labels.json")) ) train_set_mjsynth = MJSynth( train=True, img_folder='/media/cvpr/CM_22/mjsynth/mnt/ramdisk/max/90kDICT32px', label_path='/media/cvpr/CM_22/mjsynth/mnt/ramdisk/max/90kDICT32px/imlist.txt', img_transforms=T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), ) train_set_synthtext = SynthText( train=True, recognition_task=True, download=True, img_transforms=T.Resize((args.input_size, 4 * args.input_size), preserve_aspect_ratio=True), ) # 合并三个数据集 train_set = ConcatDataset([train_set_mjsynth, train_set_synthtext, custom_train_set])
方法2:修改数据集存储逻辑(不推荐)
如果一定要用extend方式,需要让所有数据集的data存储完整绝对路径而非仅图片名称,同时修改_read_sample方法直接使用该路径读取图片。但这种方式需要改动底层数据集代码,复杂度高,优先选方法1。
额外注意点
- 确保三个数据集的
__getitem__返回格式一致(图片张量+目标标签),ConcatDataset才能正常工作。 - 自定义数据集有增强操作,MJSynth和SynthText没有,若需要统一增强,可将增强逻辑移到DataLoader的
collate_fn,或统一设置到所有数据集的img_transforms中。
内容的提问来源于stack exchange,提问作者Khawar Islam
相关产品推荐
相关产品推荐

