CUDA Thrust根据指定键列表获取关联值及编译报错排查
Thrust键值查询编译错误解决方案
核心错误原因
你的代码存在4个关键问题导致编译失败:
thrust::copy_if要求传入迭代器作为输出参数,你直接传了output容器对象而非output.begin()迭代器,这是触发编译错误的直接原因- 自定义谓词不能直接持有
thrust::device_vector实例,设备端无法访问device_vector的成员方法和重载运算符,必须传入设备内存原始指针+长度 - 伪代码中键值类型写反:
keys1存储的是整数键,类型应该为long,values存储三维向量,类型应该为vec3 - 无需手动构造
zip_iterator,使用copy_if的模板(stencil)重载版本即可直接实现按键筛选值的逻辑,代码更简洁不易出错
修正后的代码实现
1. 自定义谓词
struct FindValue { const long* to_find_ptr; size_t to_find_len; __host__ __device__ explicit FindValue(const long* ptr, size_t len) : to_find_ptr(ptr), to_find_len(len) {} __host__ __device__ bool operator()(const long& key) const { for (size_t i = 0; i < to_find_len; ++i) { if (key == to_find_ptr[i]) return true; } return false; } };
2. 主逻辑调用
void correctedCode() { const int N = 7; // 修正键值类型定义 thrust::device_vector<long> keys1(N); thrust::device_vector<vec3> values(N); // 赋值操作(和你原有逻辑一致,仅修正类型对应关系) values[0] = vec3(1.01,1.01,1.0156); values[1] = vec3(1.01,1.01,1.01561); values[2] = vec3(1.02,1.52,1.02); values[3] = vec3(1.02,1.52,1.02); values[4] = vec3(1.0,1.0,1.0); values[5] = vec3(5.0,1.0,1.0); values[6] = vec3(5.0,1.5,1.0); keys1[0] = 0; keys1[1] = 1; keys1[2] = 2; keys1[3] = 5; keys1[4] = 9; keys1[5] = 19; keys1[6] = 22; // 待查询键列表,按需设置大小 thrust::device_vector<long> to_find(3); to_find[0] = 1; to_find[1] = 5; to_find[2] = 9; // 预分配输出空间,最大长度不超过原始值的数量 thrust::device_vector<vec3> output(N); // 调用stencil版本copy_if,按键筛选对应值 auto output_end = thrust::copy_if( values.begin(), values.end(), keys1.begin(), output.begin(), FindValue(thrust::raw_pointer_cast(to_find.data()), to_find.size()) ); // 裁剪输出到实际匹配的长度 output.resize(thrust::distance(output.begin(), output_end)); }
补充说明
- 如果待查询的
to_find列表长度较大,建议先对to_find排序,谓词中使用二分查找替代线性遍历,可以大幅提升查询效率 - 如果需要保留查询键的顺序、支持重复键查询,可以改用
thrust::lower_bound实现,无需遍历所有原始键值对
内容的提问来源于stack exchange,提问作者Poterie UN PETIT TOUR DE TERRE
相关产品推荐
相关产品推荐

