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

优化numpy数组加载与resize流程 降低ML训练内存占用

内存占用过高的核心原因

当前实现的内存开销主要来自三个可优化点:

  • 流程存在冗余存储:先把所有原始尺寸的npy数组全部加载存入列表,再遍历列表生成pad后的全尺寸数组,同一份数据在内存中同时存了两份,且每个pad生成的临时零矩阵都会额外占用内存
  • 数据类型未做优化:numpy加载数组若未指定dtype,浮点数默认使用float64(单值占8字节),标签默认使用int64(单值占8字节),存在大量不必要的空间浪费
  • 列表存储存在内存碎片:用list逐次append数组会产生大量内存碎片,实际占用的内存比连续存储的numpy数组高15%-30%
优化方案

按实现难度和收益排序,可按需选择:

方案1:单次遍历预分配数组(改造成本最低,内存降低70%+)

完全去掉冗余的双份存储逻辑:第一遍遍历仅读取文件元数据(路径、标签、shape),统计总样本数、全局最大高度、最大宽度后,直接预分配连续内存的最终numpy数组;第二遍逐文件加载npy,直接将pad后的数据写入预分配数组的对应位置,加载完单个文件立刻释放原始数组内存,全程不保留冗余数据。
核心优化点:选择满足精度要求的最小dtype是收益最高的手段。视觉类任务用float32(单值占4字节)即可满足精度要求,内存比float64直接减半;如果原始数据是0-255的整数类型,直接用uint8(单值占1字节),内存仅为float64的1/8。标签仅3个类别,用int8存储即可。
实现代码:

import os
import numpy as np

def load_padded_dataset(path, dtype=np.float32):
    # 第一遍遍历:仅收集元数据,不加载全量数组
    file_list = []
    labels = []
    max_h = 0
    max_w = 0
    CHANNEL = 6
    class_map = {"Class 1": 0, "Class 2": 1, "Class 3": 2}

    for folder in os.listdir(path):
        folder_path = os.path.join(path, folder)
        if not os.path.isdir(folder_path):
            continue
        label = class_map[folder]
        for file_name in os.listdir(folder_path):
            if not file_name.endswith(".npy"):
                continue
            file_path = os.path.join(folder_path, file_name)
            # 用mmap模式只读文件头获取shape,内存开销可忽略
            tmp = np.load(file_path, mmap_mode="r")
            h, w, c = tmp.shape
            assert c == CHANNEL, f"文件{file_path}通道数异常,期望{CHANNEL},实际{c}"
            max_h = max(max_h, h)
            max_w = max(max_w, w)
            file_list.append(file_path)
            labels.append(label)
            del tmp

    total_num = len(file_list)
    # 一次性预分配连续内存的最终数组,无内存碎片
    x = np.zeros((total_num, max_h, max_w, CHANNEL), dtype=dtype)
    y = np.array(labels, dtype=np.int8)

    # 第二遍遍历:逐文件加载,直接写入预分配数组
    for idx, file_path in enumerate(file_list):
        arr = np.load(file_path).astype(dtype, copy=False)
        h, w = arr.shape[:2]
        x_start = (max_w - w) // 2
        y_start = (max_h - h) // 2
        x[idx, y_start:y_start+h, x_start:x_start+w, :] = arr
        del arr

    return x, y, max_h, max_w

该方案下,原500个样本15GB的内存占用,在float32精度下可降到3GB以内,若用uint8存储可降到1GB以内。

方案2:懒加载数据集(内存占用与总样本量无关,适配超大样本量场景)

如果总样本量过大,哪怕预分配单个数组仍然无法装入内存,可采用按需加载的逻辑:初始化时仅存储文件路径、标签、全局宽高等元数据(内存开销可忽略),训练时需要取样本才加载对应npy,现场做padding,凑够一个batch输入模型后立刻释放该batch的内存。该方案下内存占用仅和batch size挂钩,哪怕几十万样本也可在低配环境运行。
以PyTorch数据集实现为例:

import os
import numpy as np
from torch.utils.data import Dataset

class Npy6ChannelDataset(Dataset):
    def __init__(self, path, dtype=np.float32):
        self.dtype = dtype
        self.file_list = []
        self.labels = []
        self.max_h = 0
        self.max_w = 0
        self.CHANNEL = 6
        class_map = {"Class 1": 0, "Class 2": 1, "Class 3": 2}

        # 初始化仅扫元数据
        for folder in os.listdir(path):
            folder_path = os.path.join(path, folder)
            if not os.path.isdir(folder_path):
                continue
            label = class_map[folder]
            for file_name in os.listdir(folder_path):
                if not file_name.endswith(".npy"):
                    continue
                file_path = os.path.join(folder_path, file_name)
                tmp = np.load(file_path, mmap_mode="r")
                h, w, c = tmp.shape
                assert c == self.CHANNEL
                self.max_h = max(self.max_h, h)
                self.max_w = max(self.max_w, w)
                self.file_list.append(file_path)
                self.labels.append(label)
                del tmp

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

    def __getitem__(self, idx):
        # 取样本时才加载单个文件,现场padding
        file_path = self.file_list[idx]
        arr = np.load(file_path).astype(self.dtype, copy=False)
        h, w = arr.shape[:2]
        pad_arr = np.zeros((self.max_h, self.max_w, self.CHANNEL), dtype=self.dtype)
        x_start = (self.max_w - w) // 2
        y_start = (self.max_h - h) // 2
        pad_arr[y_start:y_start+h, x_start:x_start+w, :] = arr
        del arr
        return pad_arr, self.labels[idx]
额外优化建议
  • 不要使用pickle存储预处理数据:pickle序列化numpy数组的效率极低,存储占用比原生npy更大,加载速度更慢,无使用价值
  • 若全局最大宽高远大于样本平均尺寸,可改用95分位的宽高作为统一尺寸,剩余5%的超大样本做下采样适配,可大幅降低单样本尺寸,内存占用线性下降
  • 加载npy不需要写入时,统一加mmap_mode="r"参数,通过操作系统内存映射按需读取数据,可进一步降低峰值内存

内容的提问来源于stack exchange,提问作者Julian Vega

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 19:30:09