---
title: GroupTopkOperation
description: "| 硬件型号 | 是否支持 |"
url: https://www.hiascend.com/document/detail/zh/canncommercial/latest/API/ascendtb/ascendtb_01_0154.html
sourcePath: /source/zh/canncommercial/900/API/ascendtb/ascendtb_01_0154.html
indexId: c2f3c45a2437262e4d49fb16d7974b170310f3e93312e9d642eb8c7194cb5d1364
---
# GroupTopkOperation

#### 产品支持情况

| 硬件型号 | 是否支持 |
| --- | --- |
| Atlas 350 加速卡 | x |
| Atlas A3 推理系列产品 / Atlas A3 训练系列产品 | √ |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | √ |
| Atlas 训练系列产品 | x |
| Atlas 推理系列产品 | x |
| Atlas 200I/500 A2 推理产品 | x |


#### 功能说明

GroupTopk算子超参数。将输入tensor0中维度1（输入tensor0有2个维度：维度0和维度1）数据分groupNum个组，每组取最大值，然后选出每组最大值中前k个，最后将非前k个组的数据全部置零。


#### 定义

```
struct GroupTopkParam {
    int32_t groupNum = 1;
    int32_t k = 0;
    enum GroupMultiFlag : uint16_t {
        UNDEFINED = 0, 
        SUM_MULTI_MAX  
    };
    uint16_t n = 1; 
    uint8_t rsv[12] = {0};
};
```


#### 参数列表

| 成员名称 | 类型 | 默认值 | 取值范围 | 是否必选 | 描述 |
| --- | --- | --- | --- | --- | --- |
| groupNum | int32\_t | 1 | [1, expert\_num] | 否 | 每个token分组数量。注：expert\_num为输入参数token维度dim\_1的值。 |
| k | int32\_t | 0 | [1, groupNum] | 是 | 选择top K专家数量。需要大于等于1。 |
| groupMultiFlag | uint16\_t | UNDEFINED | [0,1] | 否 | 枚举值，组内取值计算类型。 UNDEFINED：默认类型，每组内取最大值。 SUM\_MULTI\_MAX：每组内取n个最大值求和，需要设置参数n。 |
| n | uint16\_t | 1 | [1, expert\_num/groupNum] | 否 | 每组内取值的个数。 groupMultiFlag为1时，n需要大于0。 groupMultiFlag为0时不生效。 |
| rsv[12] | uint8\_t | {0} | \- | \- | 预留参数。 |


#### 输入

| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| token | [dim\_0, dim\_1] | float16/bf16 | ND | 输入tensor0， 二维tensor，dim\_0为token数，dim\_1为专家总数。 |
| idxArr | [1024] | int32 | ND | 输入tensor1， 一维tensor，用于辅助计算，固定长度1024，[0,1,2,...,1023]的等差序列。 |


#### 输出

| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| output | [dim\_0, dim\_1] | float16/bf16 | ND | 输出tensor，只有一个输出tensor，是对输入tensor0原地写的输出。数据类型与输入tensor0保持一致。 |


#### 约束说明

1≤expert_num≤1024，expert_num≥groupNum≥k≥1，expert_num能被groupNum整除。


#### 基础功能

- 功能概述
  将输入的各专家分数分组，每组内选取最大值，根据每组最大值大小选取topk组，其余组置零。

- 计算公式
对于每个token，首先将token的数据均匀分groupNum组，

  1. 每组取最大值；
  2. 对每组最大值降序排序；
  3. 获取前k个组的index，然后将其余的index对应的组的数据置零。
- 计算图
  groupMultiFlag取0时，计算过程下图所示：

  图1 功能示例

- 参数列表
  见参数列表，需要满足groupMultiFlag置为UNDEFINED。

- 输入| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| token | [num\_tokens, expert\_num] | float16/bf16 | ND | 二维tensor。 维度0为token数。 维度1为专家总数。 |
| idxArr | [1024] | int32 | ND | 一维tensor，用于辅助计算，固定长度1024，[0,1,2,...,1023]的等差序列。 |


- 输出| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| output | [num\_tokens, expert\_num] | float16/bf16 | ND | 是对输入token原地写的输出。数据类型与输入token保持一致。 |


- 使用示例
  1 2 3 4 // 参数构造 atb::infer::GroupTopkParam param; param.groupNum = 8; param.k = 1;


#### 组内再排序和选取前n个值功能

- 功能概述
  将输入的各专家分数分组，每组内降序排列选取前n个值求和，根据和的大小选取topk组，其余组置零。

- 计算公式
对于每个token，首先将token的数据均匀分groupNum组，

  1. 每组降序排列，取前n个值；
  2. 对这n个值求和；
  3. 各个组根据求得的和大小降序排序；
  4. 获取前k个组的index，然后将其余的index对应的组的数据置零。
n取1时对应基础功能场景。


- 参数列表
  见参数列表，需要满足groupMultiFlag置为SUM_MULTI_MAX。

- 输入| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| token | [num\_tokens, expert\_num] | float16/bf16 | ND | 二维tensor。 维度0为token数。 维度1为专家总数。 |
| idxArr | [1024] | int32 | ND | 一维tensor，用于辅助计算，固定长度1024，[0,1,2,...,1023]的等差序列。 |


- 输出| 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- |
| output | [num\_tokens, expert\_num] | float16/bf16 | ND | 是对输入tensor0原地写的输出。数据类型与输入tensor0保持一致。 |


- 使用示例
  1 2 3 4 5 6 // 参数构造 atb::infer::GroupTopkParam param; param.groupNum = 8; param.k = 1; param.groupMultiFlag = 1; param.n = 2;
