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

Python多进程池搭配NVTX标记:子进程标记未显示问题咨询

问题

在Python multiprocessing.Pool 中使用NVTX标记时,子进程调用带@nvtx.annotate注解的函数时,标记不会出现在Nsight Systems性能分析报告中,仅父进程调用时标记可见。

复现代码

import os
import time
from multiprocessing import Pool, shared_memory

import numpy as np
import nvtx

N_SAMPLES = int(1e6)
SIGNAL = np.random.randn(N_SAMPLES) + 1j * np.random.randn(N_SAMPLES)

@nvtx.annotate(color="red")
def create_shm_array(signal):
    # Store the signal in shared memory to share across processes
    shm = shared_memory.SharedMemory(create=True, size=signal.nbytes)
    shared_array = np.ndarray(signal.shape, dtype=signal.dtype, buffer=shm.buf)
    shared_array[:] = signal[:]
    return shm

def worker(shm_name):
    shm = shared_memory.SharedMemory(name=shm_name)
    sig = np.ndarray((N_SAMPLES,), dtype=complex, buffer=shm.buf)
    return expensive_op(sig)

@nvtx.annotate(color="blue")
def expensive_op(sig):
    time.sleep(2)
    return np.sum(sig)

def clean_shm(shm_name):
    shm = shared_memory.SharedMemory(name=shm_name)
    shm.close()
    shm.unlink()

if __name__ == "__main__":
    print(f"Total num_bytes: {SIGNAL.nbytes} B | {SIGNAL.nbytes / 1e9} GB")
    test = np.random.randn(10)
    expensive_op(test)
    shared_mem = create_shm_array(SIGNAL)

    with Pool(os.cpu_count()) as p:
        p.map(worker, [shared_mem.name] * 2)
    clean_shm(shared_mem.name)

现象

Nvidia Nsight Systems时间线中,父进程首次调用expensive_op时的蓝色NVTX标记清晰可见,但子进程调用该函数时无任何标记显示。

解决方法与原因分析

原因

Unix系统下multiprocessing.Pool默认用fork创建子进程,fork后的子进程会继承父进程的内存空间,但NVTX的内部跟踪状态(比如会话上下文、线程本地存储)无法在子进程中正常工作,导致Nsight抓不到子进程的NVTX标记。

可行解决方案

  1. 在子进程中显式激活NVTX上下文
    在worker函数开头添加一个空的NVTX注解,触发子进程的NVTX初始化,确保后续注解能被捕获。修改后的worker函数如下:

    def worker(shm_name):
        # 触发子进程NVTX初始化
        with nvtx.annotate(color="green", message="worker_init"):
            pass
        shm = shared_memory.SharedMemory(name=shm_name)
        sig = np.ndarray((N_SAMPLES,), dtype=complex, buffer=shm.buf)
        return expensive_op(sig)
    
  2. 改用spawn方式启动进程池
    spawn会重新启动Python解释器,重新初始化所有模块状态(包括NVTX),从根本上解决fork带来的状态继承问题。修改主进程代码:

    if __name__ == "__main__":
        import multiprocessing
        # 设置进程启动方式为spawn
        multiprocessing.set_start_method('spawn')
        
        print(f"Total num_bytes: {SIGNAL.nbytes} B | {SIGNAL.nbytes / 1e9} GB")
        test = np.random.randn(10)
        expensive_op(test)
        shared_mem = create_shm_array(SIGNAL)
    
        with Pool(os.cpu_count()) as p:
            p.map(worker, [shared_mem.name] * 2)
        clean_shm(shared_mem.name)
    

    注意:spawn会增加进程启动开销,且要求传递给子进程的对象必须可被pickle序列化,不过这里用共享内存传递数据,不会有问题。

  3. 手动注册子进程线程到NVTX
    如果需要更底层的控制,可以在子进程中调用nvtx.register_thread(需确认NVtx Python API支持),强制将线程纳入NVTX跟踪范围。不过前两种方法已经能覆盖大部分场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 13:50:22