在C++中为Eigen实现top_k:三维张量按指定维度取前k值与索引
在Eigen中对三维张量沿指定维度执行Top-K操作
针对形状为[B, D, C]的Eigen::Tensor<float, 3, Eigen::RowMajor>张量,要沿第二维度(D维度)提取前k个最大值及其索引,输出形状为[B, k, C]的结果张量,可通过以下方式实现(对应TensorFlow的tf.math.top_k功能):
实现思路
Eigen Tensor库没有内置的Top-K函数,我们可以通过维度重排+逐元素排序的方式实现:
- 调整张量维度顺序为
[B, C, D],方便对每个样本-通道组合的D个元素单独处理; - 对每个组合的D个元素,按值降序排序后取前k个,记录对应的数值和原始索引;
- 将结果整理回
[B, k, C]的目标形状。
完整代码实现
头文件依赖
#include <Eigen/Dense> #include <Eigen/Tensor> #include <vector> #include <algorithm> #include <iostream>
Top-K核心函数
// 输入:形状[B, D, C]的张量,需提取的top-k数量k // 输出:values([B,k,C])存储前k个最大值,indices([B,k,C])存储对应原始索引 void tensor_top_k(const Eigen::Tensor<float, 3, Eigen::RowMajor>& input, int k, Eigen::Tensor<float, 3, Eigen::RowMajor>& values, Eigen::Tensor<int, 3, Eigen::RowMajor>& indices) { const int B = input.dimension(0); const int D = input.dimension(1); const int C = input.dimension(2); // 处理k超过D的边界情况 k = std::min(k, D); // 重排维度为[B, C, D],便于遍历每个样本-通道的D元素 auto input_reshaped = input.swap_layout().shuffle(Eigen::array<int, 3>{0, 2, 1}); // 初始化输出张量 values.resize(Eigen::array<int, 3>{B, k, C}); indices.resize(Eigen::array<int, 3>{B, k, C}); // 遍历所有样本和通道 for (int b = 0; b < B; ++b) { for (int c = 0; c < C; ++c) { // 获取当前样本-通道对应的D元素切片 Eigen::Tensor<float, 1, Eigen::RowMajor> slice = input_reshaped.chip(b, 0).chip(c, 0); // 绑定元素值与原始索引 std::vector<std::pair<float, int>> elements; elements.reserve(D); for (int d = 0; d < D; ++d) { elements.emplace_back(slice(d), d); } // 按元素值降序排序 std::sort(elements.begin(), elements.end(), [](const auto& a, const auto& b) { return a.first > b.first; }); // 填充前k个结果到输出张量 for (int i = 0; i < k; ++i) { values(b, i, c) = elements[i].first; indices(b, i, c) = elements[i].second; } } } }
示例验证(对应你给出的TensorFlow案例)
int main() { // 构造输入张量(形状[2,3,4]) Eigen::Tensor<float, 3, Eigen::RowMajor> input(2, 3, 4); input.setValues({ {{3, 9, 8, 5}, {3, 3, 0, 1}, {5, 5, 1, 9}}, {{1, 3, 0, 9}, {8, 3, 7, 7}, {8, 0, 9, 5}} }); int k = 2; Eigen::Tensor<float, 3, Eigen::RowMajor> values; Eigen::Tensor<int, 3, Eigen::RowMajor> indices; tensor_top_k(input, k, values, indices); // 打印Values结果 std::cout << "Values:\n"; for (int b = 0; b < values.dimension(0); ++b) { std::cout << "["; for (int i = 0; i < values.dimension(1); ++i) { std::cout << "["; for (int c = 0; c < values.dimension(2); ++c) { std::cout << values(b, i, c) << (c == values.dimension(2)-1 ? "" : ", "); } std::cout << "]" << (i == values.dimension(1)-1 ? "" : ", "); } std::cout << "]\n"; } // 打印Indices结果 std::cout << "\nIndices:\n"; for (int b = 0; b < indices.dimension(0); ++b) { std::cout << "["; for (int i = 0; i < indices.dimension(1); ++i) { std::cout << "["; for (int c = 0; c < indices.dimension(2); ++c) { std::cout << indices(b, i, c) << (c == indices.dimension(2)-1 ? "" : ", "); } std::cout << "]" << (i == indices.dimension(1)-1 ? "" : ", "); } std::cout << "]\n"; } return 0; }
运行后输出结果将与你提供的TensorFlow示例完全一致。
内容的提问来源于stack exchange,提问作者MTT
相关产品推荐
相关产品推荐

