You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch超大图像数据集训练:洗牌与归一化内存溢出问题求助

大规模PyTorch数据集的内存高效洗牌与归一化方案

一、解决划分前的洗牌问题

核心思路是只对样本索引洗牌,而非直接操作全量数据,原始内存映射数据全程不加载进内存,仅通过索引定位读取样本:

  1. 生成全局索引数组:创建包含所有样本序号(0到64151)的数组,该数组体积极小(仅约250KB),可直接在内存中完成洗牌。
  2. 基于洗牌后的索引划分数据集:拆分索引数组得到训练/测试集对应的索引子集,无需触碰原始数据。
  3. 自定义Dataset类:通过内存映射读取数据,根据索引返回对应样本。

示例代码:

import numpy as np
import torch
from torch.utils.data import Dataset, Subset, DataLoader

class MemMapImageDataset(Dataset):
    def __init__(self, mmap_file_path, mean=None, std=None):
        # 内存映射加载数据集,仅建立磁盘映射,不加载全量数据
        self.data = np.memmap(
            mmap_file_path,
            dtype='float32',
            mode='r',
            shape=(64152, 3, 5, 2, 64, 144)
        )
        self.mean = mean
        self.std = std

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx):
        # 仅读取单个样本到内存
        sample = self.data[idx].copy()  # copy避免内存映射的只读限制
        if self.mean is not None and self.std is not None:
            # 应用归一化
            sample = (sample - self.mean.reshape(3, 1, 1, 1, 1)) / self.std.reshape(3, 1, 1, 1, 1)
        return torch.tensor(sample, dtype=torch.float32)

# 生成并洗牌全局索引
total_indices = np.arange(64152)
np.random.shuffle(total_indices)

# 划分训练/测试集(按8:2比例)
train_split = int(0.8 * len(total_indices))
train_indices = total_indices[:train_split]
test_indices = total_indices[train_split:]

# 初始化数据集
full_dataset = MemMapImageDataset("your_dataset_path.npy")
train_dataset = Subset(full_dataset, train_indices)
test_dataset = Subset(full_dataset, test_indices)

# 构建DataLoader,训练时开启batch级洗牌
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)

二、内存高效的归一化方案

避免全量数据加载计算统计量,采用分块遍历计算全局均值/方差,再在线应用归一化:

  1. 分块计算全局统计量:通过内存映射分批次读取数据,累加计算均值和方差,全程仅加载单个batch的数据到内存。
  2. 在线归一化:在Dataset的__getitem__中对单个样本应用预计算的均值和方差,或在collate_fn中对batch进行归一化。

示例代码(计算均值和方差):

import numpy as np

# 内存映射加载数据集
data_mmap = np.memmap(
    "your_dataset_path.npy",
    dtype='float32',
    mode='r',
    shape=(64152, 3, 5, 2, 64, 144)
)

# 初始化统计量(针对通道维度,即第1维)
mean = np.zeros(3, dtype='float64')
var = np.zeros(3, dtype='float64')
total_samples = 0
batch_size = 1024  # 根据内存容量调整批次大小

# 计算均值
for start in range(0, len(data_mmap), batch_size):
    end = min(start + batch_size, len(data_mmap))
    batch = data_mmap[start:end]
    # 沿样本、5、2、64、144维度求均值
    batch_mean = batch.mean(axis=(0, 2, 3, 4, 5))
    mean += batch_mean * (end - start)
    total_samples += (end - start)
mean /= total_samples

# 计算方差
for start in range(0, len(data_mmap), batch_size):
    end = min(start + batch_size, len(data_mmap))
    batch = data_mmap[start:end]
    # 广播均值维度,计算方差
    centered = batch - mean.reshape(1, 3, 1, 1, 1, 1)
    batch_var = (centered ** 2).mean(axis=(0, 2, 3, 4, 5))
    var += batch_var * (end - start)
var /= total_samples
std = np.sqrt(var)

# 将统计量传入数据集
full_dataset = MemMapImageDataset("your_dataset_path.npy", mean=mean, std=std)

关键注意事项

  • 洗牌仅操作索引,原始数据全程驻留磁盘,彻底避免全量加载的内存峰值。
  • 计算统计量时使用float64类型累加,避免精度损失。
  • 若使用多进程DataLoader,确保内存映射文件路径为绝对路径,避免子进程路径问题。

内容的提问来源于stack exchange,提问作者Clayton Malott

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 01:12:17