ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch数据加载器深度解析:Dataset与DataLoader核心原理与工业实战

PyTorch数据加载器深度解析:Dataset与DataLoader核心原理与工业实战 1. 项目概述为什么“加载数据集”是PyTorch学习真正的分水岭刚接触PyTorch时很多人以为写个import torch、定义个nn.Module、跑通一个loss.backward()就算入门了。我带过三十多个从零起步的工程师和研究生八成卡在第六到第七天——不是败给反向传播的链式法则而是栽在torch.utils.data.Dataset和DataLoader这两行代码上。你可能已经用torchvision.datasets.MNIST跑通了第一个手写数字识别但只要换一个本地存的CSV文件、一张没按标准命名的JPEG图、或者一个带嵌套结构的JSON标注立刻报错KeyError: image、RuntimeError: stack expects each tensor to be equal size、甚至更诡异的OSError: DataLoader worker (pid xxx) is killed by signal: Bus error。这不是你代码写错了而是你还没真正理解PyTorch数据加载机制的设计哲学它根本不是“把文件读进来”而是一套按需调度、内存隔离、多进程协同的生产级数据流水线。“PyTorch学习七——加载数据集”这个标题看似平淡实则直指整个深度学习工程落地的核心瓶颈。它背后藏着三个必须穿透的认知层第一层是语法层知道Dataset.__getitem__要返回什么、DataLoader的num_workers设多少第二层是系统层理解Python GIL如何与子进程通信、共享内存页如何避免重复拷贝、pin_memory为何能加速GPU传输第三层是工程层当你的数据集从1GB涨到1TB、从单机扩展到分布式训练、从静态图像变成实时视频流时怎么让数据加载不拖慢显卡利用率——这时候DataLoader就不再是API调用而是整个训练吞吐量的闸门。我去年帮一家工业质检公司优化YOLOv8训练流程把数据加载耗时从每batch 320ms压到47msGPU利用率从58%拉到92%核心改动就三处重写了__getitem__的缓存策略、调整了prefetch_factor、启用了persistent_workers。这些细节官方文档不会告诉你“为什么必须这么改”但实战中差1毫秒都可能让模型收敛慢两天。所以这篇内容不是教你怎么抄代码而是带你亲手拆开PyTorch数据加载器的外壳看清里面的齿轮怎么咬合。你会看到为什么Dataset必须实现__len__和__getitem__这两个魔法方法为什么DataLoader开启多进程后反而更慢collate_fn到底在哪个环节起作用pin_memoryTrue真的总能加速吗我会用真实场景——比如处理“焊接缺陷数据集V2”这种带不规则尺寸图像和稀疏标注的工业数据或者“加州房价数据集”这种混杂数值/类别/缺失值的结构化表格——一步步演示从原始文件到可训练张量的完整链路。无论你是刚写完第一个Linear层的新手还是正在调试分布式训练的算法工程师这里没有废话只有踩过坑后验证过的硬核逻辑。2. 核心设计思路PyTorch数据加载器的三层架构与选型逻辑2.1 为什么不用Pandas直接读CSV喂模型——数据加载的本质矛盾很多初学者会疑惑既然Pandas能轻松读取CSV、HDF5、Parquet为什么PyTorch还要搞一套DatasetDataLoader我拿“加州房价数据集”举个例子。假设你用pd.read_csv(housing.csv)加载后直接torch.tensor(df.values)转成张量再用torch.utils.data.TensorDataset包装——这确实能跑通但隐藏着三个致命问题第一是内存爆炸风险。该数据集有20640条记录每条含9个特征。Pandas默认用float64存储数值单条占72字节全量加载就是1.5MB但实际训练时你只需要当前batch的32条样本却提前占用了全部内存。当数据集扩大到百万级这种“全量预加载”会让8GB显存的机器直接OOM。第二是I/O阻塞瓶颈。Pandas读取是同步操作CPU必须等磁盘IO完成才能继续。而现代GPU如RTX 4090处理一个batch只需2-3ms但机械硬盘读取32条样本可能耗时15ms——GPU有90%时间在空转等数据。PyTorch的DataLoader通过多进程预取prefetch把IO和计算并行化本质是用内存换时间。第三是数据增强耦合性。工业场景中“焊接缺陷数据集V2”的图像需要做随机旋转、亮度扰动、缺陷区域mask填充这些操作必须在CPU端完成GPU不擅长图像像素级运算且要保证每次__getitem__返回的都是新增强结果。如果用Pandas预加载所有增强必须在内存里做既浪费资源又无法实现真正的随机性。所以PyTorch的设计选择非常明确Dataset负责定义“如何获取单个样本”DataLoader负责解决“如何高效供给批量样本”。前者是数据源的抽象接口后者是高性能数据管道的调度引擎。这种分离让开发者能自由组合——你可以用torchvision.datasets.ImageFolder加载标准目录结构也可以为“桥墩病害数据集”自定义Dataset解析XML标注还能用WebDataset直接流式读取网络上的tar包而DataLoader对它们一视同仁。2.2 Dataset不只是一个类而是数据契约的法律文书torch.utils.data.Dataset看似简单只强制要求实现两个方法但它实际是一份严谨的“数据契约”。我见过太多人把__getitem__写成这样def __getitem__(self, idx): img_path self.img_list[idx] image cv2.imread(img_path) # 返回BGR格式numpy数组 label self.labels[idx] return image, label # 错返回numpy数组而非tensor这段代码在小数据集上能跑但埋下三个隐患类型不一致DataLoader默认用default_collate函数堆叠张量遇到numpy数组会自动转tensor但cv2.imread返回的是uint8而模型通常期望float32导致后续归一化出错维度混乱OpenCV读图是(H,W,C)PyTorch要求(C,H,W)不转换会导致卷积核错位无错误防护idx超出范围时cv2.imread返回Nonedefault_collate堆叠None直接崩溃。正确的契约履行方式必须包含四要素确定性同一idx永远返回相同样本便于验证集复现原子性__getitem__内完成所有IO和预处理不依赖外部状态类型规范返回torch.Tensor或可被collate_fn处理的原生类型异常兜底对损坏文件、缺失标注主动抛出ValueError而非静默失败。以“声音振动信号电机数据集”为例其原始文件是.mat格式含采样率、通道数、时序信号三重信息。我的MotorSignalDataset实现会这样处理def __getitem__(self, idx): mat_file self.mat_files[idx] try: data scipy.io.loadmat(mat_file) # 提取信号矩阵确保维度为(1, T)即单通道时序 signal data[signal].reshape(1, -1) # 截断或补零至统一长度T1024 if signal.shape[1] 1024: signal np.pad(signal, ((0,0), (0, 1024-signal.shape[1]))) else: signal signal[:, :1024] # 归一化到[-1,1] signal signal / np.max(np.abs(signal) 1e-8) # 转为float32 tensor return torch.from_numpy(signal).float(), self.labels[idx] except Exception as e: raise ValueError(fFailed to load {mat_file}: {str(e)})这里每个步骤都是契约条款的具象化reshape保证维度确定性pad/slice实现长度原子性/np.max完成类型规范try-except提供异常兜底。当你把Dataset当作法律文书来写后续所有环节才不会崩塌。2.3 DataLoader参数背后的硬件博弈论DataLoader的参数表面是配置项实则是CPU、内存、磁盘、GPU四者间的资源博弈。我用一张表揭示关键参数的真实含义参数默认值实际影响工程建议batch_size1决定GPU显存占用和梯度累积步数从32开始试用nvidia-smi监控显存逐步翻倍直到OOMnum_workers0子进程数0表示主进程加载单线程CPU核心数-1但SSD上超过4个worker收益递减pin_memoryFalse是否将tensor锁页内存加速GPU传输必须True除非内存不足配合non_blockingTrue使用drop_lastFalsebatch不足时是否丢弃训练设True避免最后batch尺寸不同导致BN层异常验证设Falseprefetch_factor2每个worker预取batch数SSD设2HDD设1NVMe可设3-4persistent_workersFalseworker进程是否复用大数据集必开避免反复fork开销最关键的博弈发生在num_workers。很多人盲目设为CPU核心数结果发现训练变慢。原因在于每个worker进程启动时会复制主进程的内存空间包括已加载的模型权重若模型有500MB8个worker就额外吃掉4GB内存更糟的是Linux的fork()在内存压力大时会触发写时复制Copy-on-Write导致IO等待加剧。我在Ubuntu服务器上实测过“CWRU轴承数据集”1.2GBnum_workers0时每epoch 120snum_workers4降到85s但num_workers8反而升到98s——因为内存带宽被worker间的数据拷贝占满。另一个常被忽视的点是prefetch_factor。它的本质是“流水线缓冲区大小”。设为2意味着当GPU处理batch#0时worker#0在准备batch#1worker#1在准备batch#2。但如果磁盘IO慢如机械硬盘读取大图像buffer填不满GPU仍会等。我处理“KITTI数据集”每张图4MB时将prefetch_factor从2提到4配合persistent_workersTrue使GPU利用率从65%提升到89%。这说明参数不是孤立的必须结合你的硬件栈SSD/NVMe/RAID、数据尺寸图像分辨率/音频采样率、模型复杂度ResNet50 vs ViT动态调整。3. 实操全流程从零构建工业级数据加载器3.1 场景还原焊接缺陷数据集V2的加载挑战我们以真实工业数据集“the welding defect dataset v2”为蓝本。该数据集包含12,480张PNG图像分辨率从640×480到1920×1080不等对应XML标注文件含缺陷类型、边界框坐标5类缺陷裂纹、气孔、未熔合、夹渣、焊瘤图像质量参差部分存在运动模糊、低对比度、强反光传统做法是用ImageFolder或CocoDetection但这里行不通ImageFolder要求严格目录结构class/subclass/img.jpg而该数据集是平铺的CocoDetection依赖COCO格式JSON需手动转换XML且不支持多尺度图像直接输入更关键的是工业检测需要保持原始分辨率进行高精度定位不能简单resize到固定尺寸。因此我们必须自定义WeldingDefectDataset。整个流程分四步数据探查→路径索引→样本加载→批处理适配。第一步数据探查——用脚本代替肉眼检查先写个探查脚本避免后期踩坑import os import xml.etree.ElementTree as ET from PIL import Image import numpy as np def inspect_dataset(root_dir): img_paths [] xml_paths [] sizes [] for root, _, files in os.walk(root_dir): for f in files: if f.lower().endswith(.png): img_paths.append(os.path.join(root, f)) elif f.lower().endswith(.xml): xml_paths.append(os.path.join(root, f)) print(fFound {len(img_paths)} images, {len(xml_paths)} XML files) # 检查配对完整性 img_basenames set([os.path.splitext(p)[0] for p in img_paths]) xml_basenames set([os.path.splitext(p)[0] for p in xml_paths]) missing_xml img_basenames - xml_basenames missing_img xml_basenames - img_basenames print(fMissing XML for {len(missing_xml)} images) print(fMissing image for {len(missing_img)} XMLs) # 统计图像尺寸分布 for p in img_paths[:100]: # 取样100张 try: with Image.open(p) as img: sizes.append(img.size) except: print(fCorrupted image: {p}) sizes np.array(sizes) print(fSize range: {sizes.min(axis0)} to {sizes.max(axis0)}) print(fMean size: {sizes.mean(axis0).astype(int)}) inspect_dataset(/data/welding_v2)运行结果暴露关键问题12,480张图中137张缺失XML23张XML无对应图像尺寸跨度极大最小640×480最大1920×1080均值1280×72011张图像损坏PIL打开报错。这些发现直接决定后续设计必须加try-except容错尺寸处理不能简单resize缺失样本需在__len__中过滤。第二步构建索引——用内存换效率的底层逻辑Dataset.__init__里绝不做IO操作正确做法是预生成索引列表class WeldingDefectDataset(torch.utils.data.Dataset): def __init__(self, root_dir, transformNone, target_transformNone): self.root_dir root_dir self.transform transform self.target_transform target_transform # 预构建索引只存路径不加载数据 self.samples [] # [(img_path, xml_path, class_id), ...] self.classes [crack, porosity, lack_of_fusion, slag, weld_bead] for root, _, files in os.walk(root_dir): for f in files: if f.lower().endswith(.png): img_path os.path.join(root, f) xml_path os.path.splitext(img_path)[0] .xml if os.path.exists(xml_path): # 解析XML获取class_id try: tree ET.parse(xml_path) obj tree.find(object) if obj is not None: cls_name obj.find(name).text.strip() if cls_name in self.classes: class_id self.classes.index(cls_name) self.samples.append((img_path, xml_path, class_id)) except: continue # 跳过损坏XML print(fValid samples: {len(self.samples)}) def __len__(self): return len(self.samples)这里的关键洞察是索引构建是离线过程应在__init__一次完成而非__getitem__实时扫描。self.samples列表在初始化时就确定了所有有效样本后续__getitem__只需O(1)索引访问。我测试过对12,480个文件os.walk构建索引耗时1.2秒而每次__getitem__实时找文件平均要8ms——十万次访问就是800秒足够训练一个epoch了。第三步样本加载——处理多尺度图像的实战技巧__getitem__是性能热点必须精打细算def __getitem__(self, idx): img_path, xml_path, class_id self.samples[idx] # 1. 加载图像用OpenCV比PIL快30%且支持更多格式 try: image cv2.imread(img_path) if image is None: raise ValueError(fFailed to load {img_path}) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # BGR-RGB except Exception as e: raise ValueError(fImage load error {img_path}: {e}) # 2. 解析XML获取bbox只取第一个object简化工业场景 try: tree ET.parse(xml_path) obj tree.find(object) bbox [int(obj.find(bndbox/xmin).text), int(obj.find(bndbox/ymin).text), int(obj.find(bndbox/xmax).text), int(obj.find(bndbox/ymax).text)] except Exception as e: # 工业数据常有标注缺失此时返回全图作为bbox bbox [0, 0, image.shape[1], image.shape[0]] # 3. 应用transform重点保持原始宽高比 if self.transform: # 使用Albumentations库它原生支持bbox坐标变换 transformed self.transform(imageimage, bboxes[bbox], labels[class_id]) image transformed[image] bbox transformed[bboxes][0] if transformed[bboxes] else [0,0,0,0] # 4. 转为tensor并归一化 image torch.from_numpy(image).permute(2, 0, 1).float() / 255.0 bbox torch.tensor(bbox, dtypetorch.float32) return image, bbox, class_id这里有几个硬核技巧OpenCV替代PIL在服务器环境cv2.imread比PIL.Image.open快30%-50%尤其对PNGbbox兜底策略工业数据标注常不完整用全图坐标[0,0,W,H]保证下游模型不崩溃Albumentations集成它比torchvision.transforms更擅长处理bbox且支持HorizontalFlip、RandomBrightness等工业常用增强permute顺序cv2.imread返回(H,W,C)permute(2,0,1)转为(C,H,W)这是PyTorch卷积层的输入要求。第四步批处理适配——突破default_collate的限制default_collate只能处理同尺寸张量而我们的图像尺寸各异。解决方案是自定义collate_fndef welding_collate_fn(batch): 自定义collate对多尺度图像做padding保持原始宽高比 batch: list of (image, bbox, class_id) images, bboxes, class_ids zip(*batch) # 找出batch内最大H和W max_h max(img.shape[1] for img in images) max_w max(img.shape[2] for img in images) # padding到统一尺寸左上角对齐右侧/下侧补0 padded_images [] padded_bboxes [] for img, bbox in zip(images, bboxes): h, w img.shape[1], img.shape[2] pad_h, pad_w max_h - h, max_w - w # 使用torch.nn.functional.pad比numpy.pad更高效 padded_img F.pad(img, (0, pad_w, 0, pad_h), modeconstant, value0) padded_images.append(padded_img) # bbox坐标按比例缩放因padding不改变原始坐标 scaled_bbox bbox.clone() scaled_bbox[0] * max_w / w scaled_bbox[2] * max_w / w scaled_bbox[1] * max_h / h scaled_bbox[3] * max_h / h padded_bboxes.append(scaled_bbox) return torch.stack(padded_images, 0), \ torch.stack(padded_bboxes, 0), \ torch.tensor(class_ids) # 创建DataLoader train_loader DataLoader( WeldingDefectDataset(/data/welding_v2, transformtrain_transform), batch_size16, shuffleTrue, num_workers4, collate_fnwelding_collate_fn, pin_memoryTrue, persistent_workersTrue )这个collate_fn的精妙之处在于padding而非resize保留原始分辨率细节对微小缺陷如0.1mm裂纹检测至关重要bbox坐标动态缩放padding后图像变大bbox坐标需同比例放大否则定位偏移torch.stack替代listtorch.stack比torch.cat更高效且要求所有tensor尺寸一致。实测效果在RTX 3090上batch_size16时default_collate因尺寸不一致直接报错而此方案使数据加载耗时稳定在28ms/batchGPU利用率87%。3.2 进阶实战结构化数据集的加载范式加州房价数据集当数据是CSV而非图像时Dataset设计逻辑完全不同。以“加州房价数据集”为例其字段包括MedInc(收入中位数)、HouseAge、AveRooms、Population、AveOccup、Latitude、Longitude、MedHouseVal(目标房价)。挑战在于数值型特征需标准化类别型无但地理坐标需特殊处理缺失值处理该数据集无缺失但工业数据常有目标变量MedHouseVal需分桶做分类任务。class CaliforniaHousingDataset(torch.utils.data.Dataset): def __init__(self, csv_path, splittrain, test_size0.2, feature_colsNone, target_colMedHouseVal): self.df pd.read_csv(csv_path) self.feature_cols feature_cols or [MedInc,HouseAge,AveRooms, Population,AveOccup,Latitude,Longitude] self.target_col target_col # 划分训练/验证集固定随机种子保证可复现 np.random.seed(42) indices np.random.permutation(len(self.df)) split_idx int(len(self.df) * (1-test_size)) if split train: self.df self.df.iloc[indices[:split_idx]].reset_index(dropTrue) else: self.df self.df.iloc[indices[split_idx:]].reset_index(dropTrue) # 特征工程地理坐标转极坐标更利于模型学习 self.df[R] np.sqrt(self.df[Latitude]**2 self.df[Longitude]**2) self.df[Theta] np.arctan2(self.df[Longitude], self.df[Latitude]) # 标准化用训练集统计量验证集复用 if split train: self.scaler StandardScaler() self.features self.scaler.fit_transform(self.df[self.feature_cols]) else: # 加载训练集保存的scaler此处简化实际应pickle保存 self.features self.scaler.transform(self.df[self.feature_cols]) # 目标分桶回归转分类 self.targets pd.cut(self.df[target_col], bins5, labelsFalse).values def __len__(self): return len(self.df) def __getitem__(self, idx): x torch.from_numpy(self.features[idx]).float() y torch.tensor(self.targets[idx], dtypetorch.long) return x, y # 使用示例 train_ds CaliforniaHousingDataset(housing.csv, splittrain) val_ds CaliforniaHousingDataset(housing.csv, splitval) train_loader DataLoader(train_ds, batch_size64, shuffleTrue, num_workers2)这里体现结构化数据的三大原则划分与标准化解耦split参数控制数据切分StandardScaler在训练集拟合后应用于验证集避免数据泄露特征工程前置地理坐标转极坐标R/Theta比直接用经纬度更能表达空间关系任务适配房价是连续值但分类任务更易评估pd.cut将其分为5档labelsFalse返回整数编码。4. 常见问题与排查技巧实录那些让你熬夜的隐性陷阱4.1 “Bus error”和“Killed by signal”——内存与共享的暗战最让人抓狂的错误之一是OSError: DataLoader worker (pid xxx) is killed by signal: Bus error。这通常不是代码bug而是Linux内核的OOM Killer干的。根本原因是每个num_workers进程都复制了主进程的内存镜像当模型很大如BERT-large占1.2GB且num_workers4时仅worker就吃掉4.8GB内存加上主进程和GPU显存总内存超限触发kill。排查三步法监控内存watch -n 1 free -h观察available列是否持续下降检查worker内存ps aux --sort-%mem | head -10看哪些进程吃内存最多验证OOM Killer日志dmesg -T | grep -i killed process。解决方案降低num_workers到2-3优先保证主进程内存启用persistent_workersTrue避免worker反复fork带来的内存复制对大模型用torch.cuda.empty_cache()在__getitem__末尾释放临时GPU内存终极方案改用torch.multiprocessing.set_sharing_strategy(file_system)让worker通过文件系统共享内存而非复制。提示set_sharing_strategy必须在if __name__ __main__:块内、DataLoader创建前调用否则无效。4.2 “default_collate: batch must contain tensors”——类型不一致的隐形杀手当你自定义Dataset返回PIL.Image或numpy.ndarray时default_collate会尝试转换但常因类型不一致失败。例如# 错误示范混合返回类型 def __getitem__(self, idx): if idx % 2 0: return torch.randn(3, 224, 224) # tensor else: return np.random.randn(3, 224, 224) # numpy arraydefault_collate遇到混合类型会直接报错。但更隐蔽的是cv2.imread返回uint8torch.tensor()默认转int64而模型期望float32导致后续nn.Conv2d输入类型不匹配。快速诊断法在DataLoader迭代时打印类型for i, (x, y) in enumerate(train_loader): print(fBatch {i}: x.dtype{x.dtype}, x.shape{x.shape}) if i 2: break根治方案在__getitem__末尾强制类型转换return image.float(), label.long()对图像统一用torch.from_numpy(img).permute(2,0,1).float()/255.0对标签分类用long()回归用float()。4.3 GPU利用率低迷——数据加载成为瓶颈的证据链GPU利用率低于70%时90%是数据加载问题。判断依据有三nvidia-smi显示GPU显存已占满但Volatile GPU-Util长期50%htop中CPU核心使用率30%说明worker没饱和训练日志显示time per batch波动剧烈如20ms-200ms说明IO不稳定。针对性优化清单磁盘IO瓶颈将数据集移到NVMe SSD或启用prefetch_factor4CPU瓶颈增加num_workers但需监控htop中CPU使用率超过80%则降回内存带宽瓶颈关闭pin_memory罕见仅当RAM带宽不足时数据增强瓶颈将Albumentations的transforms.Compose移到GPU端用kornia库但需权衡CUDA内存。我曾帮一个团队诊断YOLOv8训练慢的问题nvidia-smi显示GPU利用率42%htop显示CPU使用率25%。启用torch.utils.benchmark后发现DataLoader耗时占整个batch的68%。最终解决方案是将num_workers从0改为4prefetch_factor从2提到3persistent_workersTruepin_memoryTrue。优化后GPU利用率升至89%单epoch训练时间从32分钟缩短到18分钟。4.4 分布式训练中的数据加载陷阱在torch.nn.parallel.DistributedDataParallel下DataLoader需配合DistributedSamplerfrom torch.utils.data.distributed import DistributedSampler train_sampler DistributedSampler(train_dataset, num_replicasworld_size, rankrank, shuffleTrue) train_loader DataLoader(train_dataset, batch_size32, samplertrain_sampler, num_workers4, collate_fncustom_collate)常见错误忘记设置sampler导致各GPU加载相同数据等效于batch_size×world_size但梯度更新不协同shuffleTrue与sampler冲突DistributedSampler已内置shuffleDataLoader的shuffle必须设Falsedrop_lastTrue缺失当数据总量不能被world_size×batch_size整除时最后一轮各GPU数据量不等DDP会卡死。注意DistributedSampler的num_replicas必须等于GPU总数rank是当前进程的GPU ID0到world_size-1。5. 工程进阶从单机到生产环境的加载器演进5.1 WebDataset应对TB级数据集的流式方案当数据集超过1TB如“POI数据集”含十亿级地点信息传统文件系统IO成为瓶颈。WebDataset提供基于tar存档的流式加载原理是将数万张图像打包成.tar文件DataLoader直接从tar中随机seek读取避免海量小文件的inode开销。import webdataset as wds # 构建tar文件一次性的预处理 # tar -cf welding_v2.tar --formatustar -C /data/welding_v2/ . # 流式加载 dataset wds.WebDataset(welding_v2.tar) \ .decode(wds.image_handler(pil)) \ .to_tuple(jpg;png, xml, cls) \ .map(transform_func) \ .batched(16, partialFalse) loader wds.WebLoader(dataset, num_workers8, prefetch_factor4)WebDataset的优势存储效率tar压缩比zip高且免去文件系统元数据开销加载速度SSD上顺序读tar比随机读文件快5-10倍扩展性支持S3、GCS等对象存储wds.WebDataset(pipe:aws s3 cp s3://bucket/data.tar -)直接流式读取云端数据。5.2 动态组件加载应对多模态数据的架构设计现代AI系统常需同时处理图像、文本、时序信号。“开源数据集轴承齿轮”就含振动信号.mat、红外热图.png、维修日志.txt。硬编码Dataset会失控应采用插件化设计class MultiModalDataset(torch.utils.data.Dataset): def __init__(self, config): self.modality_loaders {} for modality, cfg in config.items(): if modality image: self.modality_loaders[modality] ImageLoader(cfg) elif modality signal: self.modality_loaders[modality] SignalLoader(cfg) elif modality text: self.modality_loaders[modality] TextLoader(cfg) def __getitem__(self, idx): sample {} for modality, loader in self.modality_loaders.items(): sample[modality] loader.load(idx) return sample # config.yaml # image: {root: /data/images, transform: resize_224} # signal: {root: /data/signals, fs: 10000} # text: {root: /data/logs, tokenizer: bert-base}这种设计让数据加载器具备“动态组件加载”能力新增模态只需添加XXXLoader类无需修改主逻辑。我在医疗AI项目中用此架构接入CT影像、病理切片、电子病历上线后新增
返回列表