使用thrust::sort_by_key高效排序扁平化矩阵的行
问题:基于优先级数组原地重排扁平化矩阵的行
我将一个𝑎×𝑏维度的整数矩阵以行优先顺序存储为扁平化一维数组,现在需要根据长度为𝑎的优先级数组(可类比thrust::sort_by_key中的键数组)重新排列矩阵的行。
核心问题
- 需按优先级数组的升序对矩阵行进行排序;
- 尝试将扁平化数组转为
int[b]形式,但因b不是编译时常量,C++不支持变长数组,无法实现; - 无法承担完整复制矩阵的开销,但可使用少量额外缓冲区。
约束条件
- 矩阵以连续一维数组按行优先顺序存储;
- 优先级数组提供行的排序规则;
- 需要原地排序,仅可使用小型辅助缓冲区。
示例代码
int main() { int a = 3, b = 4; thrust::device_vector<int> matrix = { 1, 2, 3, 4, // 第0行 5, 6, 7, 8, // 第1行 9, 10, 11, 12 // 第2行 }; thrust::device_vector<int> priorities = {2, 0, 1}; // 行应重排为 [第1行, 第2行, 第0行] // 预期输出: // 5 6 7 8 // 9 10 11 12 // 1 2 3 4 }
解决方案
思路
- 先通过
sort_by_key对优先级和行索引排序,得到重排后的行顺序,仅需O(a log a)时间和O(a)的辅助空间; - 使用一行大小的临时缓冲区,通过循环交换行数据实现原地重排,避免完整复制矩阵;
- 用动态索引访问行数据,避开C++变长数组的限制。
具体实现代码
#include <thrust/device_vector.h> #include <thrust/sort.h> #include <thrust/copy.h> #include <thrust/sequence.h> #include <iostream> int main() { int a = 3, b = 4; thrust::device_vector<int> matrix = { 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12 }; thrust::device_vector<int> priorities = {2, 0, 1}; thrust::device_vector<int> row_indices(a); thrust::sequence(row_indices.begin(), row_indices.end()); // 按优先级排序行索引,得到重排后的行顺序 thrust::sort_by_key(priorities.begin(), priorities.end(), row_indices.begin()); // 仅用一行大小的临时缓冲区 thrust::device_vector<int> temp_row(b); // 原地交换行数据 for (int i = 0; i < a; ++i) { // 当前行已在正确位置,跳过 if (row_indices[i] == i) continue; int target_idx = row_indices[i]; // 保存目标行到临时缓冲区 thrust::copy(matrix.begin() + target_idx * b, matrix.begin() + (target_idx + 1) * b, temp_row.begin()); // 将当前行移到目标行位置 thrust::copy(matrix.begin() + i * b, matrix.begin() + (i + 1) * b, matrix.begin() + target_idx * b); // 将临时行放到当前行位置 thrust::copy(temp_row.begin(), temp_row.end(), matrix.begin() + i * b); // 更新索引,避免重复处理 for (int j = i + 1; j < a; ++j) { if (row_indices[j] == i) { row_indices[j] = target_idx; break; } } } // 输出验证 for (int i = 0; i < a; ++i) { for (int j = 0; j < b; ++j) { std::cout << matrix[i * b + j] << " "; } std::cout << std::endl; } return 0; }
方案说明
- 空间开销仅为O(a + b),其中a是行索引数组大小,b是临时行缓冲区大小,属于题目允许的小型辅助空间;
- 无需依赖编译时常量b,通过动态计算行起始索引访问数据,避开了C++变长数组的限制;
- 全程原地操作,没有完整复制矩阵,符合低开销要求。
内容的提问来源于stack exchange,提问作者Samiran K
相关产品推荐
相关产品推荐

