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

在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函数,我们可以通过维度重排+逐元素排序的方式实现:

  1. 调整张量维度顺序为[B, C, D],方便对每个样本-通道组合的D个元素单独处理;
  2. 对每个组合的D个元素,按值降序排序后取前k个,记录对应的数值和原始索引;
  3. 将结果整理回[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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 18:25:31