如何生成图像与掩码路径CSV并高效加载60k数据到PyTorch DataLoader?
高效加载语义分割数据到PyTorch DataLoader的方案
一、生成图像-掩码路径映射CSV文件
针对你的目录结构,图像和掩码文件名仅前缀不同,可以通过批量替换前缀的方式快速生成路径映射CSV,同时加入文件存在性检查避免无效条目:
import os import csv # 根目录路径(根据实际路径调整) root_dir = "./Segmentation" images_dir = os.path.join(root_dir, "images") masks_dir = os.path.join(root_dir, "masks") csv_save_path = os.path.join(root_dir, "data_pairs.csv") # 获取排序后的图像文件名(保证顺序对应) image_filenames = sorted(os.listdir(images_dir)) with open(csv_save_path, "w", newline="") as csv_file: writer = csv.writer(csv_file) writer.writerow(["image_path", "mask_path"]) # 写入表头 for img_name in image_filenames: # 替换前缀得到对应掩码文件名 mask_name = img_name.replace("color_left_trajectory", "color_segmentation") # 构造绝对路径(避免相对路径的歧义) img_full_path = os.path.abspath(os.path.join(images_dir, img_name)) mask_full_path = os.path.abspath(os.path.join(masks_dir, mask_name)) if os.path.exists(mask_full_path): writer.writerow([img_full_path, mask_full_path]) else: print(f"警告:未找到掩码文件 {mask_name},跳过该图像")
如果文件名匹配逻辑更复杂(比如前缀格式不固定),可以用正则表达式提取关键标识:
import re pattern = re.compile(r"color_left_trajectory_(\d+)_(\d+)\.jpg") match = pattern.match(img_name) if match: mask_name = f"color_segmentation_{match.group(1)}_{match.group(2)}.jpg"
二、基于CSV的自定义Dataset实现
编写PyTorch Dataset类读取CSV并加载数据,适配语义分割任务的掩码格式要求:
import torch from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as transforms class SegmentationDataset(Dataset): def __init__(self, csv_path, img_transform=None, mask_transform=None): self.data_pairs = [] # 读取CSV中的路径对 with open(csv_path, "r") as csv_file: reader = csv.DictReader(csv_file) self.data_pairs = [(row["image_path"], row["mask_path"]) for row in reader] self.img_transform = img_transform self.mask_transform = mask_transform def __len__(self): return len(self.data_pairs) def __getitem__(self, idx): img_path, mask_path = self.data_pairs[idx] # 加载图像为RGB格式,掩码为单通道灰度图 image = Image.open(img_path).convert("RGB") mask = Image.open(mask_path).convert("L") # 应用预处理 if self.img_transform: image = self.img_transform(image) if self.mask_transform: mask = self.mask_transform(mask) # 掩码转为LongTensor(适配CrossEntropyLoss等分类损失) mask = torch.LongTensor(mask) return image, mask
三、高效加载的核心优化建议
- 多进程加速加载:在DataLoader中设置
num_workers为CPU核心数的1-2倍,同时开启pin_memory=True(GPU训练时,减少CPU到GPU的数据传输耗时):
# 示例:定义预处理和DataLoader transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) mask_transform = transforms.Compose([ transforms.Resize((256, 256), interpolation=Image.NEAREST) # 掩码用最近邻插值避免标签模糊 ]) dataset = SegmentationDataset("Segmentation/data_pairs.csv", transform, mask_transform) dataloader = DataLoader(dataset, batch_size=16, shuffle=True, num_workers=8, pin_memory=True)
高性能数据增强:替换torchvision为
albumentations库,它支持GPU加速的增强操作,速度更快且支持更多语义分割专用的增强(比如随机裁剪、翻转时同步处理图像和掩码)。数据划分与缓存:如果需要训练/验证拆分,可以用
sklearn.model_selection.train_test_split拆分CSV数据并保存为独立的训练/验证CSV;若内存充足,可对高频访问数据做内存缓存,避免重复磁盘IO。批量处理优化:如果图像尺寸不一致,自定义
collate_fn实现动态批量调整(比如按尺寸分组、统一裁剪),避免因尺寸差异导致的加载效率下降。
内容的提问来源于stack exchange,提问作者Siddharth
相关产品推荐
相关产品推荐

