TopKFusion
=================


沿指定轴查找最大或最小 `topk` 个值的索引。当 `topk=1` 且 `largest=true` 时，该算子等价于 `ArgMax`；当 `topk=1` 且 `largest=false` 时，该算子等价于 `ArgMin`。

对于输入张量沿指定轴的每个切片 :math:`X_{slice_i}`，本算子返回其中最大（`largest=true`）或最小（`largest=false`）的 `topk` 个元素的索引：

.. math::

    Y_{i,j} = \operatorname{TopKIndex}_{j}(X_{slice_i}, largest), \quad j = 0, 1, \dots, topk-1

其中 :math:`Y_{i,j}` 表示第 :math:`j` 个结果的索引，返回结果按值从大到小（`largest=true`）或从小到大（`largest=false`）排列。

输入：
    - **input** - 输入数据地址。
    - **params** - 其他参数打包成数组。
    - **core_mask** - 核掩码（仅共享存储版本使用）。

输出：
    - **output** - 存储索引的输出张量，数据类型为 int32。
    - **output_value** - 存储找到的值，数据类型与输入相同。

.. note::
    每次调用都会同时填充 `output`（索引）和 `output_value`（值），不存在 `return_values` 开关。

.. note::
    返回的 `topk` 个结果按值有序排列：
    - 当 `largest=true` 时，按值从大到小排列；
    - 当 `largest=false` 时，按值从小到大排列。

支持平台：
    ``FT78NE``
    ``MT7004``

.. note::
    - FT78NE 支持 fp32、fp64、int32、int16、int8
    - MT7004 支持 fp16、fp32、int32、int16

**参数数组结构：**

.. code-block:: c
    :linenos:

    long long params[11];
    params[0]  = (long long)in_shape;      // 输入张量的维度信息数组
    params[1]  = (long long)in_strides;    // 输入张量的步长信息数组
    params[2]  = (long long)out_strides;   // 输出张量的步长信息数组
    params[3]  = (long long)arg_elements;  // 用于存放候选值的临时工作空间地址，大小为 axis_dim * core_num * sizeof(输入数据类型)，8 字节对齐
    params[4]  = (long long)index;         // 用于存放候选索引的临时工作空间地址，大小为 axis_dim * core_num * sizeof(int32)，8 字节对齐
    params[5]  = (long long)topk;          // 需要查找的最大/最小值的数量
    params[6]  = (long long)in_shape_size; // 输入张量的维度数 (即 in_shape 数组的长度)
    params[7]  = (long long)axis;          // 执行查找操作的轴
    params[8]  = (long long)largest;       // 是否查找最大值的标志。若为 0，则查找最小值
    params[9]  = (long long)topk_val;      // 临时缓存，大小为 topk * core_num * sizeof(输入数据类型)，8 字节对齐
    params[10] = (long long)topk_idx;      // 临时缓存，大小为 topk * core_num * sizeof(int32)，8 字节对齐

**共享存储版本:**

.. c:function:: void fp_topk_s(float *input, void *output, float *output_value, long long *params, int core_mask)
.. c:function:: void hp_topk_s(half *input, void *output, half *output_value, long long *params, int core_mask)
.. c:function:: void dp_topk_s(double *input, void *output, double *output_value, long long *params, int core_mask)
.. c:function:: void i32_topk_s(int *input, void *output, int *output_value, long long *params, int core_mask)
.. c:function:: void i16_topk_s(int16_t *input, void *output, int16_t *output_value, long long *params, int core_mask)
.. c:function:: void i8_topk_s(int8_t *input, void *output, int8_t *output_value, long long *params, int core_mask)

**C调用示例：**

.. code-block:: c
    :linenos:
    :emphasize-lines: 45

    // MT7004 示例
    #include <stdio.h>
    #include <topk.h>
    int main(int argc, char* argv[]) {
        float* input = (float*)0x81000000; // 需要初始化
        void* output = (void*)0x82000000;  // 不需要初始化，输出 int32 索引
        float* output_value = (float*)0x83000000; // 不需要初始化，输出对应值
        float* arg_elements = (float*)0x84000000; // 不需要初始化
        float* topk_val = (float *)0x89000000;
        int* topk_idx = (int *)0x8A000000;
        int* index = (int*)0x85000000; // 不需要初始化

        int core_mask = 0b1111;

        int *in_strides = (int*)0x86000000;
        int *out_strides = (int*)0x86000200;

        int in_shape_size = 4; // 最多只考虑4维
        int in_shape[4] = {4, 8, 16, 8};

        int axis = 1; // 要操作的维度，不能超过3
        int topk = 3; // 不超过 in_shape[axis]

        srand(time(0));
        // 初始化测试数据，包含各种情况
        int i;
        int in_total_elements = in_shape[0] * in_shape[1] * in_shape[2] * in_shape[3];
        for(i = 0; i < in_total_elements; i ++) {
            input[i] = (float)(rand()%100);
        }

        long long params[11];
        params[0] = (long long)in_shape;
        params[1] = (long long)in_strides;
        params[2] = (long long)out_strides;
        params[3] = (long long)arg_elements;
        params[4] = (long long)index;
        params[5] = (long long)topk;
        params[6] = (long long)in_shape_size;
        params[7] = (long long)axis;
        params[8] = (long long)1; // largest: 1 表示最大值，0 表示最小值
        params[9] = (long long)topk_val;
        params[10] = (long long)topk_idx;

        fp_topk_s(input, output, output_value, params, core_mask);
        return 0;
    }


**私有存储版本:**

.. c:function:: void fp_topk_p(float *input, void *output, float *output_value, long long *params)
.. c:function:: void hp_topk_p(half *input, void *output, half *output_value, long long *params)
.. c:function:: void dp_topk_p(double *input, void *output, double *output_value, long long *params)
.. c:function:: void i32_topk_p(int *input, void *output, int *output_value, long long *params)
.. c:function:: void i16_topk_p(int16_t *input, void *output, int16_t *output_value, long long *params)
.. c:function:: void i8_topk_p(int8_t *input, void *output, int8_t *output_value, long long *params)


**C调用示例：**

.. code-block:: c
    :linenos:
    :emphasize-lines: 43

    // MT7004 示例
    #include <stdio.h>
    #include <topk.h>
    int main(int argc, char* argv[]) {
        float* input = (float*)0x10010000; // 需要初始化
        void* output = (void*)0x10020000;  // 不需要初始化，输出 int32 索引
        float* output_value = (float*)0x10030000; // 不需要初始化，输出对应值
        float* arg_elements = (float*)0x10040000; // 不需要初始化
        float* topk_val = (float *)0x10048000;
        int* topk_idx = (int *)0x10058000;
        int* index = (int*)0x10050000; // 不需要初始化

        int *in_strides = (int*)0x1004E000;
        int *out_strides = (int*)0x1004E200;

        int in_shape_size = 4; // 最多只考虑4维
        int in_shape[4] = {4, 8, 16, 8};

        int axis = 1; // 要操作的维度，不能超过3
        int topk = 3; // 不超过 in_shape[axis]

        srand(time(0));
        // 初始化测试数据，包含各种情况
        int i;
        int in_total_elements = in_shape[0] * in_shape[1] * in_shape[2] * in_shape[3];
        for(i = 0; i < in_total_elements; i ++) {
            input[i] = (float)(rand()%100);
        }

        long long params[11];
        params[0] = (long long)in_shape;
        params[1] = (long long)in_strides;
        params[2] = (long long)out_strides;
        params[3] = (long long)arg_elements;
        params[4] = (long long)index;
        params[5] = (long long)topk;
        params[6] = (long long)in_shape_size;
        params[7] = (long long)axis;
        params[8] = (long long)1; // largest: 1 表示最大值，0 表示最小值
        params[9] = (long long)topk_val;
        params[10] = (long long)topk_idx;

        fp_topk_p(input, output, output_value, params);
        return 0;
    }