如何使用thrust::remove_if检查并成对移除数组元素块
如何用Thrust实现成对元素的批量移除
直接使用thrust::remove_if处理单个元素无法直接实现成对移除(因为它是逐元素判断的),但你可以通过以下两种可靠的方式达成需求:
方法一:索引过滤+元素收集
这种方法先筛选出需要保留的元素对索引,再批量收集对应元素,逻辑清晰且效率较高:
- 生成元素对索引序列:创建从
0到n-1的索引,每个索引对应原数组中的一对元素(Arr[2i], Arr[2i+1])。 - 过滤索引:用
thrust::remove_if移除那些满足f(a,b)=true的元素对索引,得到保留的索引列表。 - 生成元素收集索引:将每个保留的元素对索引转换为原数组中的两个元素索引(
2i和2i+1)。 - 收集结果:用
thrust::gather将原数组中对应索引的元素复制到结果数组。
示例代码:
#include <thrust/device_vector.h> #include <thrust/remove.h> #include <thrust/gather.h> #include <thrust/iterator/counting_iterator.h> // 假设f是已定义的布尔函数,注意参数类型匹配数组的uint64_t bool f(uint64_t a, uint64_t b); int main() { // 初始化原数组,长度为2n thrust::device_vector<uint64_t> Arr = {a1, b1, a2, b2, ..., an, bn}; int n = Arr.size() / 2; // 创建元素对的索引序列:0, 1, ..., n-1 auto idx_begin = thrust::make_counting_iterator(0); auto idx_end = idx_begin + n; // 移除需要删除的元素对索引(f返回true时移除该对) auto new_idx_end = thrust::remove_if(idx_begin, idx_end, [&](int i) { uint64_t a = Arr[2 * i]; uint64_t b = Arr[2 * i + 1]; return f(a, b); }); // 计算保留的元素对数量 int keep_count = new_idx_end - idx_begin; // 生成要收集的原数组元素索引 thrust::device_vector<int> gather_indices(2 * keep_count); thrust::transform(idx_begin, new_idx_end, gather_indices.begin(), [](int i) { return 2 * i; }); thrust::transform(idx_begin, new_idx_end, gather_indices.begin() + keep_count, [](int i) { return 2 * i + 1; }); // 收集元素到结果数组 thrust::device_vector<uint64_t> result(2 * keep_count); thrust::gather(gather_indices.begin(), gather_indices.end(), Arr.begin(), result.begin()); return 0; }
方法二:Zip迭代器+掩码复制
这种方法先通过Zip迭代器将元素对打包,生成掩码标记保留的对,再批量复制:
- 打包元素对:用
thrust::make_zip_iterator将原数组的偶数位和奇数位迭代器打包,得到每个元素为(ai, bi)的tuple迭代器。 - 生成保留掩码:用
thrust::transform遍历所有元素对,生成掩码数组(true表示保留该对,false表示移除)。 - 批量复制保留元素:通过掩码过滤,将需要保留的元素对复制到结果数组。
示例代码(设备端批量优化版):
#include <thrust/device_vector.h> #include <thrust/zip_iterator.h> #include <thrust/transform.h> #include <thrust/scatter.h> #include <thrust/iterator/counting_iterator.h> bool f(uint64_t a, uint64_t b); int main() { thrust::device_vector<uint64_t> Arr = {a1, b1, a2, b2, ..., an, bn}; int n = Arr.size() / 2; // 打包元素对为tuple迭代器 auto zip_begin = thrust::make_zip_iterator(thrust::make_tuple(Arr.begin(), Arr.begin() + 1)); auto zip_end = zip_begin + n; // 生成掩码:true表示保留该对 thrust::device_vector<bool> mask(n); thrust::transform(zip_begin, zip_end, mask.begin(), [&](const thrust::tuple<uint64_t, uint64_t>& pair) { return !f(thrust::get<0>(pair), thrust::get<1>(pair)); }); // 计算保留的元素对数量 int keep_count = thrust::count(mask.begin(), mask.end(), true); // 生成目标位置索引 thrust::device_vector<int> target_indices(2 * keep_count); thrust::device_vector<int> pair_positions(n); thrust::copy(thrust::make_counting_iterator(0), thrust::make_counting_iterator(n), pair_positions.begin()); // 过滤得到保留的元素对索引 auto kept_pair_begin = thrust::copy_if(pair_positions.begin(), pair_positions.end(), mask.begin(), pair_positions.begin()); // 生成原数组元素的目标索引 thrust::transform(pair_positions.begin(), kept_pair_begin, target_indices.begin(), [](int i) { return 2 * i; }); thrust::transform(pair_positions.begin(), kept_pair_begin, target_indices.begin() + keep_count, [](int i) { return 2 * i + 1; }); // 生成原数组元素的源索引 thrust::device_vector<int> source_indices(2 * n); thrust::copy(thrust::make_counting_iterator(0), thrust::make_counting_iterator(2*n), source_indices.begin()); // 初始化结果数组 thrust::device_vector<uint64_t> result(2 * keep_count, 0); // 散射复制保留的元素 thrust::scatter_if(source_indices.begin(), source_indices.end(), target_indices.begin(), [&](int idx) { return mask[idx / 2]; }, result.begin()); return 0; }
关键注意事项
- 不要直接用
thrust::remove_if处理原数组的单个元素:单个元素的谓词无法关联到其所属的元素对,会导致只移除对中的一个元素,破坏数组的成对结构。 - 优先选择方法一:逻辑更简洁,避免复杂的迭代器操作,性能也更稳定,适合处理大型数组。
内容的提问来源于stack exchange,提问作者Mojtaba Valizadeh
相关产品推荐
相关产品推荐

