PyTorch加载大量图像时内存不足问题求助
问题描述
使用PyTorch加载大量图像时,当图像数量接近50000张,内核会崩溃,推测是内存限制导致。以下是20000张图像测试的代码及内存占用情况:
import os import psutil import cv2 import numpy as np import torch print(f"Before starting to loop: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB") X_data = [] y_data = [] for path in paths: img = cv2.imread(path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) X_data.append(np.array(img/255, dtype=np.uint8)) print(f"Before convert to numpy: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB") X_data = np.array(X_data) print(f"Before shuffle: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB") shuffle_index = np.random.permutation(X_data.shape[0]) X_data = X_data[shuffle_index] print(f"Before Convert to tensor: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB") X_data = torch.Tensor(X_data).view(-1, 3, 128, 128) print(f"Before save: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB") torch.save(X_data, f"X_data.pt") print(f"After save: {psutil.Process(os.getpid()).memory_info().rss / 1024 ** 3} GB")
内存输出:
Before starting to loop: 0.26 GB Before convert to numpy: 1.29 GB Before shuffle: 2.28 GB Before Convert to tensor: 2.28 GB Before save: 5.22 GB After save: 4.14 GB
疑问:
- 代码是否存在低效操作?尝试过跳过中间步骤,但
torch.cat和numpy.append速度过慢。 - 是否推荐将数据按批次存储,训练时按需加载?找不到相关入门指南。
- 疑惑50000张1281283的图像不应引发此类问题。
解决方案
一、代码中的低效/错误点
预处理逻辑错误
你将img/255的浮点结果强制转为uint8,这会把0-1之间的浮点值直接截断为0或1,完全丢失图像的灰度信息,同时后续转PyTorch张量时又会转回浮点型,既浪费计算资源又破坏数据。正确的做法是要么保留uint8格式(训练时再归一化),要么直接转为float32存储归一化后的值。内存冗余操作
- 用列表
X_data存储单个图像数组后再转numpy大数组,这个过程中内存会同时保留列表中所有小数组和新生成的大数组两份数据,直到原列表被垃圾回收,导致内存临时翻倍(从1.29GB涨到2.28GB)。 - 将numpy数组(
uint8)转为torch.Tensor时,默认会转为float32类型,内存占用直接变为原来的4倍(从2.28GB涨到5.22GB),这是内存激增的核心原因之一。
- 用列表
二、50000张图像的内存计算
按1281283的尺寸计算:
- 单张
uint8图像:128*128*3*1 = 49152字节 - 50000张
uint8总内存:50000*49152 ≈ 2.29GB - 转为
float32后总内存:2.29*4 ≈ 9.16GB
加上Python、PyTorch本身的内存开销,以及数据处理过程中的临时内存占用,很容易超过系统可用内存,导致内核崩溃。
三、优化方案
1. 修复预处理逻辑
如果要预存数据,推荐保留uint8格式减少内存占用,训练时再做归一化:
# 修改循环内的预处理 img = cv2.imread(path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) X_data.append(img) # 直接保留uint8格式,无需除以255 # 训练时再归一化 X_tensor = torch.tensor(X_data, dtype=torch.float32).view(-1,3,128,128) / 255.0
2. 按需加载(推荐)
完全不需要一次性加载所有数据,用PyTorch的Dataset和DataLoader实现按需加载,内存只需要存储一个批次的数据:
from torch.utils.data import Dataset, DataLoader class ImageDataset(Dataset): def __init__(self, paths): self.paths = paths def __len__(self): return len(self.paths) def __getitem__(self, idx): path = self.paths[idx] img = cv2.imread(path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # 这里可以直接做归一化并转张量 img_tensor = torch.tensor(img, dtype=torch.float32).permute(2,0,1) / 255.0 # 假设y_data是标签,这里可以根据路径获取标签 label = ... # 替换为你的标签获取逻辑 return img_tensor, label # 初始化数据集和加载器 dataset = ImageDataset(paths) dataloader = DataLoader(dataset, batch_size=32, shuffle=True) # 训练时迭代加载器 for imgs, labels in dataloader: # 执行训练步骤 ...
这种方式下,内存只需要存储32张图像的张量(约3212812834 ≈ 6MB),完全不会出现内存溢出问题。
3. 预存数据的优化
如果一定要预存数据,可以用numpy的memmap或者直接存储为PyTorch张量,避免内存冗余:
- 用
numpy.memmap创建磁盘映射数组,无需一次性加载到内存:
# 创建memmap数组 shape = (len(paths), 128, 128, 3) X_memmap = np.memmap('X_data.npy', dtype=np.uint8, mode='w+', shape=shape) # 逐个写入图像 for i, path in enumerate(paths): img = cv2.imread(path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) X_memmap[i] = img # 使用时加载 X_memmap = np.memmap('X_data.npy', dtype=np.uint8, mode='r', shape=shape) X_tensor = torch.tensor(X_memmap, dtype=torch.float32).view(-1,3,128,128) / 255.0
四、总结
- 你的代码核心问题是预处理错误和内存冗余操作,导致内存占用远超预期。
- 按需加载是处理大数据集的标准方案,完全不需要一次性加载所有数据,
Dataset+DataLoader是PyTorch的基础用法,逻辑简单且高效。 - 50000张图像的内存需求远超你预期的原因是数据类型转换带来的内存膨胀,优化数据类型后可以大幅降低内存占用。
内容的提问来源于stack exchange,提问作者JakobVinkas
相关产品推荐
相关产品推荐

