如何将整个数据集一次性加载到GPU以加速模型训练
全量数据集加载到GPU显存的实现方案
先算下你这套数据集的显存占用,确认可行性:
按float32精度计算,1550张3×112×112图像的总大小为1550 * 3 * 112 * 112 * 4Byte ≈ 222MB,就算叠加标签、训练缓存,占用也远低于常规GPU显存容量,完全可以一次性加载,彻底消除数据读取IO开销。
方案1:自定义预加载Dataset(兼容原有ImageFolder逻辑)
这是最适配你现有代码的方案,不需要调整原本的数据集目录结构,只需要替换原来的ImageFolder实例即可。
核心逻辑是初始化阶段就把所有样本完成确定性预处理,直接存入GPU显存,训练迭代时直接从显存读取数据,完全跳过磁盘读取、CPU到GPU拷贝的流程。
import torch from torch.utils.data import Dataset from torchvision.datasets import ImageFolder class GPUPreloadDataset(Dataset): def __init__(self, data_root, deterministic_transform=None, train_aug=None): # 复用ImageFolder的目录解析、标签映射逻辑,不用改你现有的数据集文件结构 self.base_dataset = ImageFolder(data_root, transform=deterministic_transform) self.sample_count = len(self.base_dataset) self.train_aug = train_aug # 预分配GPU显存空间,避免逐样本拼接的额外开销 sample_img, _ = self.base_dataset[0] self.img_cache = torch.zeros( (self.sample_count, *sample_img.shape), dtype=torch.float32, device="cuda" ) self.label_cache = torch.zeros( (self.sample_count,), dtype=torch.long, device="cuda" ) # 一次性完成所有样本的确定性预处理,加载到显存 for idx in range(self.sample_count): img, label = self.base_dataset[idx] self.img_cache[idx] = img.to("cuda") self.label_cache[idx] = torch.tensor(label, dtype=torch.long, device="cuda") def __len__(self): return self.sample_count def __getitem__(self, idx): img = self.img_cache[idx] label = self.label_cache[idx] # 训练时随机增强直接在显存上执行,速度极快 if self.train_aug is not None: img = self.train_aug(img) return img, label
注意:不要把随机翻转、随机擦除、随机裁剪这类随机增强逻辑放到deterministic_transform里,否则预加载后所有epoch的增强结果固定,会失去数据增强的效果。把这类随机操作传入train_aug参数即可,torchvision 0.15以上版本原生支持GPU张量的增强运算,不会带来明显开销。
方案2:TensorDataset快速实现(代码更简洁)
如果你的预处理逻辑简单,不需要复杂的自定义逻辑,可以直接预加载所有张量后用PyTorch内置的TensorDataset构建数据集:
import torch from torchvision.datasets import ImageFolder from torch.utils.data import TensorDataset, DataLoader # 加载所有样本,执行确定性预处理 base_set = ImageFolder("你的数据集根目录", transform=确定性预处理操作) img_list = [] label_list = [] for img, label in base_set: img_list.append(img) label_list.append(label) # 拼接为批量张量,直接搬到GPU all_imgs = torch.stack(img_list).to("cuda") all_labels = torch.tensor(label_list, dtype=torch.long).to("cuda") gpu_dataset = TensorDataset(all_imgs, all_labels)
DataLoader配置注意点
因为数据已经常驻GPU显存,不需要多进程读取、内存锁页这类针对CPU加载数据的优化,配置不当反而会增加额外开销:
train_loader = DataLoader( gpu_dataset, batch_size=你设置的批大小, shuffle=True, num_workers=0, # 必须设为0,不需要子进程做数据读取 pin_memory=False # 数据已经在GPU上,关闭锁页内存 )
这套配置下,数据读取环节的耗时会降到几乎可以忽略,训练速度瓶颈会完全回到模型本身的计算上,相比原来逐次从磁盘/内存读取的方式,训练效率会有非常明显的提升。
内容的提问来源于stack exchange,提问作者Twain
相关产品推荐
相关产品推荐

