使用xtensor将NumPy布尔索引逻辑转为C++代码遇问题求助
问题:xtensor实现NumPy布尔索引等价操作
问题背景
将Python NumPy代码转换为C++ xtensor代码时,遇到布尔索引操作的维度匹配问题:
Python 原代码
boxes = decoded_boxes[keep, :]
keep:维度为[19248]的布尔数组decoded_boxes:维度为[19248, 4]的数组- 预期输出
boxes:维度为[24, 4]的数组
尝试的C++代码(不符合预期)
xt::xarray<float> b1 = xt::filter(decoded_boxes, keep); std::cout << "b1 shape = " << xt::adapt(b1.shape()) << std::endl;
问题:输出维度为{24},而非预期的[24,4]
解决方案
xt::filter会将匹配元素扁平化输出,无法保留原数组的多维结构。要实现与NumPy布尔索引完全等价的行为,需使用**xt::boolean_mask**并指定过滤轴:
正确的C++代码
// 对第0轴(第一维)应用布尔掩码,保留第二维结构 xt::xarray<float> boxes = xt::boolean_mask(decoded_boxes, keep, 0); std::cout << "boxes shape = " << xt::adapt(boxes.shape()) << std::endl;
说明
xt::boolean_mask的第三个参数用于指定应用掩码的轴(此处为第0轴,对应NumPy中decoded_boxes[keep, :]的第一维过滤)- 执行后输出的
boxes维度为[24,4],与Python代码的行为完全一致
内容的提问来源于stack exchange,提问作者ABD
相关产品推荐
相关产品推荐

