Python多进程+PyTorch张量训练GNN时的内存峰值问题排查
我想用Python multiprocessing模块生成包含numpy数组的数据集,将其转换为torch张量用于GNN训练,并在每个epoch后替换数据集的一定比例。但运行下方示例代码时,每次迭代后内存都会出现急剧峰值,最终脚本因OOM-kill事件崩溃。
不同场景下的内存表现:
- 同时使用多进程+torch张量转换:内存峰值极高,最终OOM
- 不转换为torch张量:内存占用平稳,无明显尖峰
- 用循环替代多进程:内存尖峰明显降低
- 仅用单进程:内存占用最平稳
代码示例:
import numpy as np from multiprocessing import Pool, cpu_count from torch_geometric.data import Data import torch device = 'cuda' if torch.cuda.is_available() else 'cpu' def generate_buffer(): repeated_args = [arguments] * buffer_size # create batches in parallel: with Pool(processes = cpu_count()) as pool: buffer = pool.starmap(generate_batch, repeated_args) # flatten the buffer: buffer = [item for sublist in buffer for item in sublist] # convert list of numpy arrays to torch Data object containing torch GPU tensors for i in range(len(buffer)): X = torch.tensor(buffer[i][0], dtype = torch.float, device = device) edge_index = torch.tensor(buffer[i][1], dtype = torch.long, device = device) edge_attr = torch.tensor(buffer[i][2], dtype = torch.float, device = device) y = torch.tensor(buffer[i][3], dtype = torch.float, device = device) buffer[i] = Data(x=X, edge_index=edge_index, edge_attr=edge_attr, y=y) return buffer def update_buffer(buffer): # delete the first entries of the buffer: del buffer[: (replacements_per_iteration * batch_size * len(error_rate))] # create a list of repeated arguments for all processes: repeated_args = [arguments] * replacements_per_iteration # create batches in parallel: with Pool(processes = cpu_count()) as pool: new_data = pool.starmap(generate_batch, repeated_args) # flatten the data: new_data = [item for sublist in new_data for item in sublist] # convert list of numpy arrays to torch Data object containing torch GPU tensors for i in range(len(new_data)): X = torch.tensor(new_data[i][0], dtype = torch.float, device = device) edge_index = torch.tensor(new_data[i][1], dtype = torch.long, device = device) edge_attr = torch.tensor(new_data[i][2], dtype = torch.float, device = device) y = torch.tensor(new_data[i][3], dtype = torch.float, device = device) new_data[i] = Data(x=X, edge_index=edge_index, edge_attr=edge_attr, y=y) # append to buffer: buffer.extend(new_data) del new_data return buffer def generate_batch(arguments): batch = [] # need to create a different seed in every thread: np.random.seed() for _ in range(batch_size): graph = generate_sample(arguments) batch.append(graph) return batch def generate_sample(arguments): # Generating a sample (syndrome measurement of a rotated surface # code cycle (a quantum error correction scheme)) and mapping # it to a graph representation using basic numpy operations return [X, edge_index, edge_attr, y] if __name__== '__main__': buffer = generate_buffer() for i in range(num_iterations): # update the buffer: buffer = update_buffer(buffer)
内存尖峰主要由以下几个核心问题导致:
多进程数据传递的内存副本冗余:
multiprocessing的子进程生成numpy数组后,会将完整的数据副本传回主进程。主进程转换为torch GPU张量时,不会自动清理这些原始numpy数组——这就造成CPU内存同时承载两份数据:原始numpy数组,以及GPU张量对应的CPU端临时数据(即使张量已移至GPU,numpy数组的内存仍未释放)。每次迭代都会新增一批冗余数据,直到Python垃圾回收触发,期间就会出现显著的内存尖峰。张量转换的时机导致内存叠加:当前代码是在主进程接收所有子进程的numpy数据后,才逐个转换为GPU张量。这个过程中,主进程需要先把所有新生成的numpy数据全部加载到CPU内存,再逐一执行转换操作,相当于短时间内CPU内存同时承载了全部新numpy数据和正在转换的张量,直接推高内存峰值。
重复创建进程池的额外开销:每次
update_buffer都新建一个Pool,进程的创建与销毁本身会带来额外内存开销,且子进程的资源可能无法完全即时释放,叠加数据传递的内存占用,进一步加剧了内存波动。垃圾回收延迟引发的内存累积:虽然代码中使用
del删除旧数据,但Python的垃圾回收是异步触发的,旧数据的内存不会立即被释放。新数据添加时,新旧数据的内存会短暂共存,多次迭代后内存累积,最终引发OOM崩溃。
内容的提问来源于stack exchange,提问作者Moritz Lange

