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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 18:05:47