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

如何使用thrust::remove_if检查并成对移除数组元素块

如何用Thrust实现成对元素的批量移除

直接使用thrust::remove_if处理单个元素无法直接实现成对移除(因为它是逐元素判断的),但你可以通过以下两种可靠的方式达成需求:

方法一:索引过滤+元素收集

这种方法先筛选出需要保留的元素对索引,再批量收集对应元素,逻辑清晰且效率较高:

  1. 生成元素对索引序列:创建从0到n-1的索引,每个索引对应原数组中的一对元素(Arr[2i], Arr[2i+1])。
  2. 过滤索引:用thrust::remove_if移除那些满足f(a,b)=true的元素对索引,得到保留的索引列表。
  3. 生成元素收集索引:将每个保留的元素对索引转换为原数组中的两个元素索引(2i和2i+1)。
  4. 收集结果:用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迭代器将元素对打包,生成掩码标记保留的对,再批量复制:

  1. 打包元素对:用thrust::make_zip_iterator将原数组的偶数位和奇数位迭代器打包,得到每个元素为(ai, bi)的tuple迭代器。
  2. 生成保留掩码:用thrust::transform遍历所有元素对,生成掩码数组(true表示保留该对,false表示移除)。
  3. 批量复制保留元素:通过掩码过滤,将需要保留的元素对复制到结果数组。

示例代码(设备端批量优化版):

#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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 09:25:26