使用pybind11与Arrow时出现函数参数不兼容问题求助
问题描述
基于pybind11与Arrow开发Python扩展模块,代码可正常编译、链接并导入Python环境,但调用print_table和aggregate_trades函数传入pyarrow Table对象时触发TypeError,提示函数参数不兼容。
相关代码
#include <pybind11/pybind11.h> #include <arrow/stl.h> #include <arrow/table.h> #include <string> #include "Trade.h" #include "TradeAggregator.h" #include <arrow/python/pyarrow.h> #include <iostream> #include <Python.h> namespace py = pybind11; using std::to_string; void print_table(PyObject *py_table) { // convert pyobject to table auto status = arrow::py::unwrap_table(py_table); if (!status.ok()) { std::cout << "Error converting pyarrow table to arrow table" << std::endl; return; } std::shared_ptr<arrow::Table> table = status.ValueOrDie(); std::cout << "Table has " << table->num_rows() << " rows" << std::endl; } // Wrapper for return to pyarrow. void aggregate_trades(PyObject *py_table) { auto status = arrow::py::unwrap_table(py_table); // ..... } PYBIND11_MODULE(mkt_data, m) { arrow::py::import_pyarrow(); m.doc() = "Market Data processing plugin "; py::class_<MktData::Trade>(m, "Trade") .def(py::init<MktData::event_id_t, MktData::timestamp_t, double, double, bool, MktData::event_id_t, MktData::event_id_t>()) .def_readonly("trade_id", &MktData::Trade::trade_id) .def_readonly("timestamp", &MktData::Trade::timestamp) .def_readonly("price", &MktData::Trade::price) .def_readonly("quantity", &MktData::Trade::quantity) .def_readonly("consideration", &MktData::Trade::consideration) .def_readonly("is_buy", &MktData::Trade::is_buy) .def_readonly("first_id", &MktData::Trade::first_id) .def_readonly("last_id", &MktData::Trade::last_id) .def("__repr__", [](const MktData::Trade &t) { return "<Trade: trade_id=" + to_string(t.trade_id) + ", timestamp=" + to_string(t.timestamp) + ", price=" + to_string(t.price) + ", quantity=" + to_string(t.quantity) + ", consideration=" + to_string(t.consideration) + ", is_buy=" + ((t.is_buy) ? "True" : "False") + ", first_id=" + to_string(t.first_id) + ", last_id=" + to_string(t.last_id) + ">"; }); ; m.def("print_table", &print_table); m.def("aggregate_trades", &aggregate_trades, "Aggregate the trades in the arrow tables."); }
报错信息
Input In [1], in <cell line: 7>() 5 df = pd.read_csv('ADAUSDT-aggTrades-2022-09-01.zip') 6 tbl = pa.Table.from_pandas(df) ----> 7 mkt_data.print_table(tbl) TypeError: print_table(): incompatible function arguments. The following argument types are supported: 1. (arg0: _object) -> None
CMakeLists.txt配置
cmake_minimum_required(VERSION 3.14) if(${CMAKE_VERSION} VERSION_LESS 3.24) cmake_policy(VERSION ${CMAKE_MAJOR_VERSION}.${CMAKE_MINOR_VERSION}) else() cmake_policy(VERSION 3.24) endif() project(MarketDataProcessing VERSION 1.0 DESCRIPTION "Preprocessing Market Data" LANGUAGES CXX) # GoogleTest requires at least C++14 set(CMAKE_CXX_STANDARD 17) # option(BUILD_PYTHON_MODULE "Build a mkt_data python module" ON) include(FindPkgConfig) find_package(Arrow REQUIRED) add_library(mdprocessing SHARED include/Trade.h include/TradeAggregator.h include/utils.h src/Trade.cpp src/TradeAggregator.cpp ) target_include_directories(mdprocessing PUBLIC ${CMAKE_CURRENT_SOURCE_DIR}/include ) target_link_libraries(mdprocessing arrow_shared) # add_subdirectory(tests) ## Build Python module find_package(pybind11 REQUIRED) #add_library(hello_world hello_world.cpp) set(PYBIND11_PYTHON_VERSION "3.9") pybind11_add_module(mkt_data src/mkt_data_wrapper.cpp) target_include_directories(mkt_data PUBLIC /home/ruihong/.python_venvs/learning/lib/python3.9/site-packages/pyarrow/include) target_link_directories(mkt_data PUBLIC /home/ruihong/.python_venvs/learning/lib/python3.9/site-packages/pyarrow) target_link_libraries(mkt_data PRIVATE mdprocessing arrow_python)
解决方案
核心原因
pybind11无法自动将pyarrow Table对象适配到PyObject*类型的参数签名,因为PyObject*是最底层的Python C API类型,pybind11不会为其自动做类型转换,需要显式声明兼容的参数类型或使用pyarrow提供的绑定工具。
具体修改步骤
1. 调整C++函数参数与绑定代码
推荐直接使用std::shared_ptr<arrow::Table>作为函数参数,让pybind11通过arrow的内置绑定自动完成Python对象到C++对象的转换,代码更简洁且类型安全:
#include <pybind11/pybind11.h> #include <arrow/stl.h> #include <arrow/table.h> #include <string> #include "Trade.h" #include "TradeAggregator.h" #include <arrow/python/pyarrow.h> #include <iostream> #include <Python.h> namespace py = pybind11; using std::to_string; // 直接使用arrow::Table的智能指针作为参数 void print_table(std::shared_ptr<arrow::Table> table) { std::cout << "Table has " << table->num_rows() << " rows" << std::endl; } // aggregate_trades同理修改参数类型 void aggregate_trades(std::shared_ptr<arrow::Table> table) { // 你的聚合逻辑实现 } PYBIND11_MODULE(mkt_data, m) { // 必须初始化pyarrow的C API绑定 arrow::py::import_pyarrow(); m.doc() = "Market Data processing plugin "; // 保留Trade类的绑定代码不变 py::class_<MktData::Trade>(m, "Trade") .def(py::init<MktData::event_id_t, MktData::timestamp_t, double, double, bool, MktData::event_id_t, MktData::event_id_t>()) .def_readonly("trade_id", &MktData::Trade::trade_id) .def_readonly("timestamp", &MktData::Trade::timestamp) .def_readonly("price", &MktData::Trade::price) .def_readonly("quantity", &MktData::Trade::quantity) .def_readonly("consideration", &MktData::Trade::consideration) .def_readonly("is_buy", &MktData::Trade::is_buy) .def_readonly("first_id", &MktData::Trade::first_id) .def_readonly("last_id", &MktData::Trade::last_id) .def("__repr__", [](const MktData::Trade &t) { return "<Trade: trade_id=" + to_string(t.trade_id) + ", timestamp=" + to_string(t.timestamp) + ", price=" + to_string(t.price) + ", quantity=" + to_string(t.quantity) + ", consideration=" + to_string(t.consideration) + ", is_buy=" + ((t.is_buy) ? "True" : "False") + ", first_id=" + to_string(t.first_id) + ", last_id=" + to_string(t.last_id) + ">"; }); // 绑定函数,pybind11自动处理类型转换 m.def("print_table", &print_table, "Print arrow table row count"); m.def("aggregate_trades", &aggregate_trades, "Aggregate the trades in the arrow tables."); }
如果坚持使用PyObject*,可以将参数改为py::object(pybind11的封装类型),再在函数内转换:
void print_table(py::object py_table) { auto status = arrow::py::unwrap_table(py_table.ptr()); if (!status.ok()) { std::cout << "Error converting pyarrow table to arrow table" << std::endl; return; } std::shared_ptr<arrow::Table> table = status.ValueOrDie(); std::cout << "Table has " << table->num_rows() << " rows" << std::endl; } // 绑定时直接使用该函数即可,pybind11会适配py::object类型 m.def("print_table", &print_table);
2. 优化CMake配置(避免硬编码路径)
使用官方提供的ArrowPython包查找工具,替代手动指定pyarrow的头文件和库路径,提升兼容性:
## Build Python module find_package(pybind11 REQUIRED) find_package(ArrowPython REQUIRED) # 官方pyarrow包查找 set(PYBIND11_PYTHON_VERSION "3.9") pybind11_add_module(mkt_data src/mkt_data_wrapper.cpp) target_include_directories(mkt_data PUBLIC ${CMAKE_CURRENT_SOURCE_DIR}/include ${ArrowPython_INCLUDE_DIRS} # 使用官方变量引入头文件 ) target_link_libraries(mkt_data PRIVATE mdprocessing ArrowPython::ArrowPython # 使用官方目标链接pyarrow库 )
3. 验证修改
重新编译安装扩展模块后,在Python环境中测试:
import pandas as pd import pyarrow as pa import mkt_data df = pd.read_csv('ADAUSDT-aggTrades-2022-09-01.zip') tbl = pa.Table.from_pandas(df) mkt_data.print_table(tbl) # 应正常输出表格行数
额外注意事项
arrow::py::import_pyarrow()必须在模块初始化时调用,确保pyarrow的C API被正确加载。- 确保编译时使用的pyarrow版本与Python环境中的pyarrow版本完全一致,避免版本不兼容导致的转换失败。
内容的提问来源于stack exchange,提问作者Rehon
相关产品推荐
相关产品推荐

