如何在Thrust中查找向量a大于向量b对应元素的索引
解决Thrust中逐元素比较两向量并筛选索引的问题
要找出Thrust device vector a中元素大于对应位置b元素的索引,核心是利用Thrust支持多输入的算法结合二元谓词实现逐元素比较,以下是具体实现方案:
方法步骤
- 生成索引序列:用
thrust::counting_iterator生成从0开始的连续索引,对应向量的每个位置。 - 定义二元比较谓词:创建函数对象,判断a中第i个元素是否大于b中第i个元素。
- 筛选符合条件的索引:使用
thrust::copy_if,结合索引序列和比较谓词,将满足a[i] > b[i]的索引复制到结果向量中。
代码示例
#include <thrust/device_vector.h> #include <thrust/copy.h> #include <thrust/iterator/counting_iterator.h> #include <iostream> // 二元谓词:判断a对应位置元素是否大于b struct greater_than_predicate { const thrust::device_vector<int>& a; const thrust::device_vector<int>& b; greater_than_predicate(const thrust::device_vector<int>& a_, const thrust::device_vector<int>& b_) : a(a_), b(b_) {} __host__ __device__ bool operator()(int idx) const { return a[idx] > b[idx]; } }; int main() { // 示例输入 thrust::device_vector<int> a = {1, 3, 5, 6, 9}; thrust::device_vector<int> b = {2, 1, 4, 7, 8}; thrust::device_vector<int> result; // 筛选满足条件的索引 thrust::copy_if( thrust::make_counting_iterator(0), thrust::make_counting_iterator(a.size()), std::back_inserter(result), greater_than_predicate(a, b) ); // 输出结果 std::cout << "满足a[i] > b[i]的索引:"; for (int idx : result) { std::cout << idx << " "; } std::cout << std::endl; // 输出:1 2 4 return 0; }
关键说明
- 你之前对
transform的输入限制存在误解:Thrust的transform有多重重载,支持2个及以上输入迭代器。比如可以先用thrust::transform(a.begin(), a.end(), b.begin(), output.begin(), thrust::greater<int>())生成布尔向量标记符合条件的位置,再结合索引筛选,但上述copy_if方案更直接高效。 - 谓词中的
__host__ __device__修饰符确保函数可在主机和设备端执行,符合Thrust的执行要求。 - 若需判断
a[i] < b[i],只需将谓词中的>替换为<即可。
内容的提问来源于stack exchange,提问作者nyaki
相关产品推荐
相关产品推荐

