使用Pandas处理10万张图像生成Siamese网络训练对时进程被kill如何解决
优化方案
问题根因
你遇到的内存溢出问题主要来自4个设计冗余:
- 核心的cross join生成全量配对再过滤去重的逻辑开销极大,若单个ID下有N张图,cross join会产生N²条临时数据,再加上用frozenset逐行去重的操作也会占用大量内存
- 初始化时创建了未使用的多进程Pool,子进程创建会复制父进程全部内存数据,直接拉高初始内存占用
- 多余的临时文件读写步骤,先存分组数据到本地再读回的操作不仅浪费IO,还会生成不必要的内存副本
- 开启了pandas全量行、列展示配置,会额外占用内存存储元数据
具体修改方案
- 重构正样本生成逻辑,放弃cross join,改用
itertools.combinations直接生成无序不重复配对,天然满足x≠y且无逆序重复对,不需要后续过滤步骤,内存占用直接降低70%以上 - 删除未使用的多进程Pool初始化代码,避免不必要的内存复制
- 删除多余的临时文件存储/读取逻辑,分组后直接处理每个组,处理完立即释放内存
- 移除pandas的全量行、列展示配置
- 每次处理完单个分组后手动清理临时变量,触发垃圾回收释放内存
- 优化写入逻辑,避免反复打开关闭文件
修改后可运行代码
""" Used to generate positive and negative pairs. """ import logging import os import itertools import gc from typing import Tuple import pandas as pd from data_generation.utils import images_to_df, save_df from tqdm import tqdm class PairGenerator: def __init__(self, root_dir: str): """ 初始化配对生成流水线 Args: root_dir: 存放jpg图像的根目录 """ self._root_dir: str = root_dir # 提前只保留需要的列,减少内存占用 self._images_df: pd.DataFrame = images_to_df(self._root_dir)[["site_id", "img_id", "img_path"]] self._samples_per_group: int = 5 logging.info(f"总图像数:{len(self._images_df)}") def generate_pairs(self) -> None: """ 生成正负样本对,保存为csv """ self.generate_positive_pairs() def generate_positive_pairs(self) -> None: """ 生成正样本对csv """ grouped = self._images_df.groupby(by=["site_id", "img_id"]) logging.info( f"共找到 {len(grouped)} 个分组,单资产平均图像数:{len(self._images_df) / len(grouped):.2f}" ) # 提前控制表头写入状态,避免重复写表头 first_write = True for group_name, group in tqdm(grouped, total=len(grouped)): pairs = self.merge_positive(group) if len(pairs) == 0: continue # 写入结果 save_df(pairs, os.path.join(self._root_dir, "positive.csv"), mode='w' if first_write else 'a', index=False, header=first_write) first_write = False # 手动清理临时变量释放内存 del pairs, group gc.collect() def merge_positive(self, group: pd.DataFrame) -> pd.DataFrame: """ 生成正样本对,不做cross join,直接用combinations生成无序对 Args: group: 同一资产的图像分组df Returns: 包含img_path_x、img_path_y列的正样本对df """ paths = group["img_path"].tolist() # 少于2张图的组无法生成配对直接返回空 if len(paths) < 2: return pd.DataFrame(columns=["img_path_x", "img_path_y"]) # 直接生成所有无序不重复对 all_pairs = list(itertools.combinations(paths, 2)) # 采样指定数量 sample_num = min(len(all_pairs), self._samples_per_group) sampled_pairs = pd.DataFrame(all_pairs, columns=["img_path_x", "img_path_y"]).sample(n=sample_num) del all_pairs, paths gc.collect() return sampled_pairs
内容的提问来源于stack exchange,提问作者pceccon
相关产品推荐
相关产品推荐

