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

使用pybind11修改结构化numpy数组无效的问题排查

问题:pybind11传递结构化numpy数组后,C++内的修改无法同步到Python

相关代码

C++结构体定义

/* Ray structure for a single ray */
struct RTC_ALIGN(16) RTCRay
{
  float org_x;        // x coordinate of ray origin
  float org_y;        // y coordinate of ray origin
  float org_z;        // z coordinate of ray origin
  float tnear;        // start of ray segment

  float dir_x;        // x coordinate of ray direction
  float dir_y;        // y coordinate of ray direction
  float dir_z;        // z coordinate of ray direction
  float time;         // time of this ray for motion blur

  float tfar;         // end of ray segment (set to hit distance)
  unsigned int mask;  // ray mask
  unsigned int id;    // ray ID
  unsigned int flags; // ray flags
};

/* Hit structure for a single ray */
struct RTC_ALIGN(16) RTCHit
{
  float Ng_x;          // x coordinate of geometry normal
  float Ng_y;          // y coordinate of geometry normal
  float Ng_z;          // z coordinate of geometry normal

  float u;             // barycentric u coordinate of hit
  float v;             // barycentric v coordinate of hit

  unsigned int primID; // primitive ID
  unsigned int geomID; // geometry ID
  unsigned int instID[1]; // instance ID
  unsigned int instPrimID[1]; // instance primitive ID
};

/* Combined ray/hit structure for a single ray */
struct RTCRayHit
{
  struct RTCRay ray;
  struct RTCHit hit;
};

pybind11类型绑定

PYBIND11_NUMPY_DTYPE(RTCRay, org_x, org_y, org_z, tnear, dir_x, dir_y, dir_z, time, tfar, mask, id, flags);
PYBIND11_NUMPY_DTYPE(RTCHit, Ng_x, Ng_y, Ng_z, u, v, primID, geomID, instID, instPrimID);
PYBIND11_NUMPY_DTYPE(RTCRayHit, ray, hit);

C++处理函数

void ray_intersect(py::array_t<RTCRayHit, py::array::c_style>& rays) {
    auto buf = rays.request();
    RTCRayHit* ptr = static_cast<RTCRayHit*>(buf.ptr);
    std::cout << ptr[0].ray.org_x << std::endl;
    ptr->ray.org_x = 2.0;
    ptr[0].ray.org_x = 2.0;
    std::cout << ptr[0].ray.org_x << std::endl;
}

Python调用代码

import numpy as np

RTC_MAX_INSTANCE_LEVEL_COUNT = 1

dt_ray = np.dtype([
    ('org_x', 'f4'), ('org_y', 'f4'), ('org_z', 'f4'), ('tnear', 'f4'),
    ('dir_x', 'f4'), ('dir_y', 'f4'), ('dir_z', 'f4'), ('time', 'f4'),
    ('tfar', 'f4'), ('mask', 'u4'), ('id', 'u4'), ('flags', 'u4')
])

# Define the RTCHit structure
dt_hit = np.dtype([
    ('Ng_x', 'f4'), ('Ng_y', 'f4'), ('Ng_z', 'f4'),
    ('u', 'f4'), ('v', 'f4'),
    ('primID', 'u4'), ('geomID', 'u4'),
    ('instID', ('u4', (RTC_MAX_INSTANCE_LEVEL_COUNT,))),
    ('instPrimID', ('u4', (RTC_MAX_INSTANCE_LEVEL_COUNT,)))
])

# Define the RTCRayHit structure
dt_rayhit = np.dtype([
    ('ray', dt_ray),
    ('hit', dt_hit)
])

rayhit = np.zeros(1, dtype=dt_rayhit)
rayhit[0]["ray"]["org_x"] = 1.0
print(rayhit)
ray_intersect(rayhit)
print(rayhit)

运行输出

[((1., 0., 0., 0., 0., 0., 0., 0., 0., 0, 0, 0), (0., 0., 0., 0., 0., 0, 0, [0], [0]))]
1
2
[((1., 0., 0., 0., 0., 0., 0., 0., 0., 0, 0, 0), (0., 0., 0., 0., 0., 0, 0, [0], [0]))]

问题分析

函数内打印并修改了ptr[0].ray.org_x的值(从1变为2),但回到Python后数组内容未更新,核心原因是C++结构体与Python dtype的内存布局不匹配,导致pybind11无法直接映射内存,而是创建了拷贝:

  1. C中的RTCRay和RTCHit使用RTC_ALIGN(16)强制16字节对齐,会在结构体末尾添加填充字节;但Python定义的dtype没有考虑对齐,内存大小与C结构体不一致。
  2. RTCHit中的instID和instPrimID是C风格固定大小数组,pybind11的PYBIND11_NUMPY_DTYPE对这类数组的处理,与Python中('u4', (1,))的多维数组布局不兼容。

解决方案

1. 匹配C++结构体的内存对齐

先在C++中打印结构体的实际大小:

std::cout << "RTCRay size: " << sizeof(RTCRay) << std::endl;
std::cout << "RTCHit size: " << sizeof(RTCHit) << std::endl;
std::cout << "RTCRayHit size: " << sizeof(RTCRayHit) << std::endl;

根据输出的大小,在Python的dtype中添加填充字段。例如若RTCRay实际大小为64字节(原字段总大小48字节),则添加16字节填充:

dt_ray = np.dtype([
    ('org_x', 'f4'), ('org_y', 'f4'), ('org_z', 'f4'), ('tnear', 'f4'),
    ('dir_x', 'f4'), ('dir_y', 'f4'), ('dir_z', 'f4'), ('time', 'f4'),
    ('tfar', 'f4'), ('mask', 'u4'), ('id', 'u4'), ('flags', 'u4'),
    ('_pad_rtcray', 'V16')  # 填充16字节匹配对齐要求
])

同理处理RTCHit的填充字段。

2. 正确绑定C风格数组

对于大小为1的C风格数组,可以直接修改C++结构体定义,将数组改为单个变量:

struct RTC_ALIGN(16) RTCHit
{
    // ... 其他字段
    unsigned int instID;     // 替换原instID[1]
    unsigned int instPrimID; // 替换原instPrimID[1]
};

之后同步修改Python的dt_hit定义,去掉数组维度:

dt_hit = np.dtype([
    # ... 其他字段
    ('instID', 'u4'),
    ('instPrimID', 'u4'),
    # ... 填充字段
])

3. 确保数组是C连续的

创建Python数组时显式指定C连续内存布局:

rayhit = np.zeros(1, dtype=dt_rayhit, order='C')

验证修改

调整后,C++中的指针将直接指向Python数组的内存,修改操作会同步到Python端。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 02:50:56