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

pybind11虚函数重写在OpenMP多线程调用时挂起的解决方法

问题:OpenMP多线程调用pybind11重写的虚函数导致程序挂起

在使用pybind11封装C++虚函数并通过Python重写后,当OpenMP多线程循环内调用这些重写方法时,程序会挂起。原以为PYBIND11_OVERRIDE宏会自动获取GIL保证线程安全,但实际存在疏漏。

重现代码

C++封装代码

#include "pybind11/pybind11.h"
#include "pybind11/functional.h"

namespace py = pybind11;
using namespace pybind11::literals;

#include <cstdio>

class A {
public:
  A()                               { printf("A::A()\n"); }    
  virtual ~A()                      { printf("A::~A()\n"); }   
  virtual void void_func() const    { printf("A::void_func()\n"); }
  virtual int int_func(int x) const { printf("A::int_func(%d)\n", x); return x + 1; }
};

void do_threaded_stuff(const A& a) {
  int sum = 0;
#pragma omp parallel for
  for (auto i = 0u; i < 10u; ++i) {
    a.void_func();
#pragma omp critical
    {
      sum = a.int_func(sum);
    }
  }
  printf("Final sum: %d\n", sum);
}

//------------------------------------------------------------------------------
// Trampoline class for A
//------------------------------------------------------------------------------
class PYB11TrampolineA: public A {
public:
  using A::A;
  virtual void void_func() const override {
    PYBIND11_OVERRIDE(void, A, void_func);
  }
  virtual int int_func(int x) const override {
    PYBIND11_OVERRIDE(int, A, int_func, x);
  }
};

//------------------------------------------------------------------------------
// Make the module
//------------------------------------------------------------------------------
PYBIND11_MODULE(virtual_override_thread, m) {
  py::class_<A, PYB11TrampolineA> obj(m, "A");
  obj.def(py::init<>());
  obj.def("void_func", (void (A::*)() const) &A::void_func);
  obj.def("int_func", (int (A::*)(int) const) &A::int_func);

  m.def("do_threaded_stuff", (void (*)(const A&)) &do_threaded_stuff, "a"_a);
}

Python调用代码

from virtual_override_thread import *

class B(A):
    def __init__(self):
        A.__init__(self)

    def void_func(self):
        print("B::void_func")

    def int_func(self, x):
        print("B::int_func({})".format(x))
        return x + 10

a = A()
do_threaded_stuff(a)    # 正常运行

b = B()
do_threaded_stuff(b)    # OMP_NUM_THREADS>1时挂起

原因分析

PYBIND11_OVERRIDE确实会尝试获取GIL,但问题出在OpenMP创建的线程没有被Python的线程管理系统注册。Python的GIL机制要求每个线程都先完成线程状态的初始化,否则在获取GIL时可能会陷入死锁:

  • 主线程持有GIL进入OpenMP并行区域,启动多个子线程
  • 子线程调用Python重写的方法时,尝试获取GIL,但由于未注册线程状态,GIL的获取逻辑会阻塞
  • 同时主线程可能在等待子线程完成,形成循环等待导致挂起

解决方案

需要在OpenMP线程首次调用Python代码前,手动初始化Python线程状态,调用完成后清理状态。可以通过修改trampoline类的方法,在PYBIND11_OVERRIDE前后添加线程状态管理:

修改后的Trampoline类代码

class PYB11TrampolineA: public A {
public:
  using A::A;
  virtual void void_func() const override {
    // 初始化线程状态并获取GIL
    py::gil_scoped_acquire acquire;
    PYBIND11_OVERRIDE(void, A, void_func);
    // 离开作用域时自动释放GIL并清理线程状态
  }
  virtual int int_func(int x) const override {
    py::gil_scoped_acquire acquire;
    return PYBIND11_OVERRIDE(int, A, int_func, x);
  }
};

关键说明

  • py::gil_scoped_acquire会自动完成:
    1. 检查当前线程是否已注册Python线程状态,未注册则初始化
    2. 获取GIL
    3. 离开作用域时自动释放GIL,若线程是首次初始化,还会清理线程状态
  • 不需要在OpenMP区域额外处理,所有调用Python重写方法的路径都通过trampoline类的方法统一管理线程状态,确保线程安全

额外优化

如果OpenMP线程会频繁调用Python方法,可以避免每次都初始化/销毁线程状态,改用py::gil_scoped_release在不需要Python时释放GIL,需要时再获取:

void do_threaded_stuff(const A& a) {
  // 主线程进入并行区域前释放GIL,减少竞争
  py::gil_scoped_release release;
  int sum = 0;
#pragma omp parallel for
  for (auto i = 0u; i < 10u; ++i) {
    a.void_func();
#pragma omp critical
    {
      sum = a.int_func(sum);
    }
  }
  printf("Final sum: %d\n", sum);
}

这样主线程在OpenMP并行区域运行时不持有GIL,能提升多线程执行效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 16:30:20