在CUDA Thrust中融合两次归约操作的实现方法问询
单次Thrust内核调用完成嵌套归约的实现方案
答案是肯定的——你可以通过自定义迭代器和融合操作,在单次Thrust内核调用里完成reduce_by_key加reduce(甚至两次reduce_by_key)的操作,既避免中间结果占用过多内存,还能减少内核启动开销提升速度。
针对你举的「扁平化矩阵每行最大值的最小值」例子,核心思路就是把两次归约逻辑融合:让reduce_by_key(按行分组求最大值)输出的每个结果,直接作为后续reduce(求最小值)的输入,全程不存储中间的所有行最大值。
具体实现思路
你提到的transform_output_iterator确实是关键,搭配一个能维护状态的自定义functor,就能把两次归约串起来:
- 写一个functor,用来接收
reduce_by_key输出的每个行最大值,同时实时更新当前的全局最小值。 - 用
transform_output_iterator把这个functor包装成输出迭代器,传给reduce_by_key——这样每次生成一个行最大值,就会立刻被functor处理,不会把所有中间结果写到内存里。 - 最后从functor里取出最终的最小值就行。
代码示例
#include <thrust/device_vector.h> #include <thrust/reduce.h> #include <thrust/iterator/transform_output_iterator.h> #include <thrust/iterator/counting_iterator.h> #include <iostream> #include <climits> // 自定义functor:接收每个行最大值,实时维护当前最小值 struct MinAccumulator { int current_min; __host__ __device__ MinAccumulator() : current_min(INT_MAX) {} // 每次拿到一个行最大值,就更新最小值 __host__ __device__ void operator()(int value) { if (value < current_min) { current_min = value; } } }; int main() { // 扁平化矩阵:3行4列 thrust::device_vector<int> mat = { 3, 1, 4, 1, 5, 9, 2, 6, 5, 3, 5, 8 }; const int rows = 3; const int cols = 4; // 生成分组键:每行的元素对应同一个键(0,0,0,0,1,1,1,1,2,2,2,2) thrust::device_vector<int> keys(mat.size()); thrust::transform( thrust::make_counting_iterator(0), thrust::make_counting_iterator(mat.size()), keys.begin(), [cols] __device__(int idx) { return idx / cols; } ); MinAccumulator accum; // 用transform_output_iterator包装accum,作为reduce_by_key的输出迭代器 // 用discard_iterator丢弃不需要的原始输出,节省内存 auto output_iter = thrust::make_transform_output_iterator( thrust::make_discard_iterator(), [&accum] __device__(int val) { accum(val); return val; } ); // 执行reduce_by_key:按行分组求最大值,结果直接传给accum处理 thrust::reduce_by_key( keys.begin(), keys.end(), mat.begin(), thrust::make_discard_iterator(), // 分组键的输出不需要存储 output_iter, thrust::equal_to<int>(), // 键的比较逻辑 thrust::maximum<int>() // 分组内的归约逻辑(求最大值) ); // 最终结果就是accum维护的最小值 std::cout << "每行最大值中的最小值:" << accum.current_min << std::endl; // 这里输出应为4(行0最大值4,行1最大值9,行2最大值8,最小值是4) return 0; }
关键细节说明
thrust::make_discard_iterator()用来丢弃不需要的输出(比如分组键的结果),完全不用为这些数据分配内存。transform_output_iterator的操作是在reduce_by_key的内核内部同步执行的,整个过程只启动一次内核,没有额外的开销。- functor的状态是在主机端初始化,然后设备端更新的,这里因为
reduce_by_key是同步操作,所以主机端可以直接读取最终状态。
扩展到两次reduce_by_key的场景
如果你的实际需求是两次reduce_by_key(比如先按A分组归约,再按B分组归约),思路类似:只需要让自定义functor同时跟踪第二次分组的键和当前分组的归约状态,当键发生变化时处理结果即可。不过这种情况要注意设备端的线程安全,必要时可以用thrust::atomic或者借助Thrust的内置分组逻辑来保证正确性。
内容的提问来源于stack exchange,提问作者Frimousse
相关产品推荐
相关产品推荐

