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
相关产品推荐
相关产品推荐

