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

Python 3D CNN模型训练数据采集阶段内存耗尽问题排查

3D CNN训练数据采集内存耗尽问题排查与解决

问题背景

我正在开发一个实现3D CNN模型的大型个人项目,采集训练数据时出现内存耗尽的情况(数据集包含约1000个分辨率为480p的小视频)。当前做法是遍历文件夹中的所有视频,每120帧提取一帧,归一化后将该帧加入列表。

相关代码

ControllerClass实现

@dataclass
class ControllerClass:
    __ann_model: VectorModelANN
    __3d_cnn_model: Model3DCNN
    __ann_data: TrainerANNData
    __3d_cnn_data: Trainer3DCNNData


    @benchmark
    def gather_training_data(self, trainer_enum: TrainerEnum) -> Union[list, list]:
        x_training_data = []
        y_training_data = []
        counter_files = 1

        if trainer_enum.value == TrainerEnum.CNN_3D.value:
            trainer = TrainingDataGenerator3DCNN(
                vid='',
                cnn_3d_width=self.__3d_cnn_data.width_3d_cnn,
                cnn_3d_height=self.__3d_cnn_data.height_3d_cnn,
            )
        elif ...  # Not important code checker for this question

        for filename in os.listdir(os.getcwd()  + "/training data"):         
            print(f"INFO:GATHER_DATA_{trainer_enum.value}: Processing file {counter_files} with the name {filename}")
            counter_files+=1

            trainer.update_new_vid(filename)

            x_train_slice = trainer.generate_data()

            x_training_data.append(
                x_train_slice
            )
            y_training_data.append(
                transfor_file_name_into_int(filename)
            )
            print(f"INFO:GATHER_DATA: {sys.getsizeof(x_training_data)}")
        return x_training_data, y_training_data

TrainingDataGenerator3DCNN实现

class TrainingDataGenerator3DCNN:
    def __init__(
        self, 
        vid,
        cnn_3d_width,
        cnn_3d_height,
    ) -> None:
        self.__vid = vid
        self.__3d_cnn_width = cnn_3d_width
        self.__3d_cnn_height = cnn_3d_height


    def normalize(self, img: Any) -> Any:
        img = cv2.resize(img, (self.__3d_cnn_width, self.__3d_cnn_height), interpolation = cv2.INTER_AREA)
        img = cv2.GaussianBlur(img, (1,1), 0)
        img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
        
        img = img / 255

        return img


    def generate_data(self) -> list:
        final_3d_data = []

        if type(self.__vid) is str:
            self.__vid = cv2.VideoCapture("./training data/" + self.__vid)
        else:
            self.__vid = cv2.VideoCapture(self.__vid)

        count_frame = 0
        while self.__vid.isOpened():
            success, image = self.__vid.read()

            if not success:
                break

            if count_frame % 120 == 0:
                image = self.normalize(image)
                final_3d_data.append(image)
                count_frame = 0
            else:
                count_frame += 1

        self.__vid.release()
        cv2.destroyAllWindows()
        return final_3d_data


    def update_new_vid(self, filename: str) -> None:
        self.__vid = filename

尺寸常量

WIDTH_3D_CNN = 480
HEIGHT_3D_CNN = 360

内存打印日志

INFO:GATHER_DATA_CNN_3D: Processing file 56 with the name abs_51.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 57 with the name abs_52.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 58 with the name abs_53.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 59 with the name abs_54.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 60 with the name abs_55.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 61 with the name abs_56.mp4
INFO:GATHER_DATA: 568
INFO:GATHER_DATA_CNN_3D: Processing file 62 with the name abs_57.mp4
INFO:GATHER_DATA: 568

疑问

我怀疑存在内存泄漏,但找不到泄漏点。我配备了32GB内存,从未想过会不够用。请问我是否遗漏了什么?是否应该将视频分辨率调整为更低的规格(比如240p)?


问题根源与解决方案

1. 内存统计的误区

你用sys.getsizeof(x_training_data)打印的只是列表容器本身的大小,完全没有统计列表中存储的numpy帧数据的实际内存占用。以当前480*360的单通道float64帧为例,单帧内存为480*360*8字节=1,382,400字节≈1.3MB,如果每个视频提取10帧,1000个视频的总内存就达到1000*10*1.3MB≈12.7GB,再加上Python的内存开销,32GB被占满是很正常的。

2. 潜在的OpenCV资源泄漏

在generate_data方法中,当self.__vid已经是VideoCapture对象时,你又重新创建了一个新的cv2.VideoCapture(self.__vid),旧的对象没有被及时释放,累积下来会占用额外内存。修复代码如下:

def generate_data(self) -> list:
    final_3d_data = []
    # 先释放已存在的VideoCapture对象
    if isinstance(self.__vid, cv2.VideoCapture):
        self.__vid.release()
    # 重新创建新的VideoCapture
    self.__vid = cv2.VideoCapture("./training data/" + self.__vid)

    count_frame = 0
    while self.__vid.isOpened():
        success, image = self.__vid.read()
        if not success:
            break
        if count_frame % 120 == 0:
            image = self.normalize(image)
            final_3d_data.append(image)
            count_frame = 0
        else:
            count_frame += 1

    self.__vid.release()
    cv2.destroyAllWindows()
    # 重置为字符串类型,避免后续出错
    self.__vid = ""
    return final_3d_data

3. 避免一次性加载全部数据

这是解决内存耗尽的核心方案:

  • 使用生成器:不要把所有数据存在列表中,改用生成器逐批次返回数据,配合框架的批量训练接口(如Keras的fit_generator或PyTorch的DataLoader)。
  • 离线存储数据:将提取的帧数据保存为.npy、HDF5或TFRecord格式,训练时再逐批次从磁盘加载,内存仅需承载当前批次的数据。

4. 优化帧数据的内存占用

  • 降低分辨率:如果模型精度允许,将分辨率降到240p(240*180),单帧内存会直接降到原来的1/4。
  • 降低数据类型:归一化后的默认是float64,转成float32可以减半内存占用,代码修改为:
    img = (img / 255).astype(np.float32)
    

5. 正确统计内存占用

要统计实际内存使用,需要遍历列表中的所有帧数据:

import numpy as np

def calculate_total_memory(data_list):
    total = 0
    for item in data_list:
        if isinstance(item, list):
            total += calculate_total_memory(item)
        elif isinstance(item, np.ndarray):
            total += item.nbytes
    return total

# 在采集流程中打印实际内存
print(f"INFO:GATHER_DATA: Total memory used: {calculate_total_memory(x_training_data) / (1024**3):.2f} GB")

内容的提问来源于stack exchange,提问作者Cosmin Leonte

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 11:57:02