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会自动完成:- 检查当前线程是否已注册Python线程状态,未注册则初始化
- 获取GIL
- 离开作用域时自动释放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
相关产品推荐
相关产品推荐

