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

#### 产品支持情况

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


#### 功能说明

融合了MLA场景下PagedAttention输入数据处理的全过程，包括从隐状态输入开始经过rmsnorm、反量化、matmul、rope、reshapeAndCache的一系列计算。

图1 计算流程图


#### 定义

```
struct MlaPreprocessParam{
    uint32_t wdqDim = 0;
    uint32_t qRopeDim = 0;
    uint32_t kRopeDim = 0;
    float epsilon = 1e-5;
    int32_t qRotaryCoeff = 2;
    int32_t kRotaryCoeff = 2;
    bool transposeWdq = true;
    bool transposeWuq = true;
    bool transposeWuk = true;
    enum CacheMode : uint8_t {
        KVCACHE = 0,
        KROPE_CTKV,
        INT8_NZCACHE,
        NZCACHE
    };
    CacheMode cacheMode = KVCACHE;
    enum QuantMode : uint16_t {
        PER_TENSOR_QUANT_ASYMM = 0,
        PER_TOKEN_QUANT_SYMM,
        PER_TOKEN_QUANT_ASYMM,
        UNQUANT
    };
    QuantMode quantMode = PER_TENSOR_QUANT_ASYMM;
    enum BackendType : int {
        BACKEND_TYPE_ATB = 0,  
        BACKEND_TYPE_ACLNN,    
        BACKEND_TYPE_MAX,      
    };
    BackendType backendType = BACKEND_TYPE_ATB;
    uint8_t rsv[34] = {0};
};
```


#### 参数列表

| 成员名称 | 类型 | 默认值 | 取值范围 | 是否必选 | 描述 |
| --- | --- | --- | --- | --- | --- |
| wdqDim | uint32\_t | 0 | [0] | 否 | 经过matmul后拆分的dim大小。 |
| qRopeDim | uint32\_t | 0 | [0] | 否 | q传入rope的dim大小。 |
| kRopeDim | uint32\_t | 0 | [0] | 否 | k传入rope的dim大小。 |
| epsilon | float | 1e\-5 | [1e\-5] | 否 | 加在分母上防止除0。 |
| qRotaryCoeff | int32\_t | 2 | 2 | 否 | q旋转系数。 |
| kRotaryCoeff | int32\_t | 2 | 2 | 否 | k旋转系数。 |
| transposeWdq | bool | true | (true) | 否 | wdq是否转置。 |
| transposeWuq | bool | true | (true) | 否 | wuq是否转置。 |
| transposeWuk | bool | true | (true) | 否 | wuk是否转置。 |
| cacheMode | uint8\_t | 0 | [0,3] | 否 | 指定cache的类型。 |
| quantMode | uint16\_t | 0 | [0,3] | 否 | 指定RmsNorm量化的类型。 |
| rsv[34] | uint8\_t | {0} | [0] | 否 | 预留字段。 |
| 当前只有cacheMode和quantMode生效。 | 当前只有cacheMode和quantMode生效。 | 当前只有cacheMode和quantMode生效。 | 当前只有cacheMode和quantMode生效。 | 当前只有cacheMode和quantMode生效。 | 当前只有cacheMode和quantMode生效。 |


上表中类型为自定义类型的，其描述如下：

cacheMode：表示输入query和kcache的类型，其具体取值如下。

- KVCACHE：kcache和q均经过拼接后输出。
- KROPE_CTKV：输入输出的kvCache拆分为krope和ctkv，q拆分为qrope和qnope。
- INT8_NZCACHE：krope和ctkv转为NZ格式输出，ctkv和qnope经过per_head静态对称量化为int8类型。
- NZCACHE：krope和ctkv转为NZ格式输出。

quantMode：表示RmsNorm量化类型，其具体取值如下。

- PER_TENSOR_QUANT_ASYMM：per_tensor静态非对称量化，默认量化类型。
- PER_TOKEN_QUANT_SYMM：per_token动态对称量化，未实现。
- PER_TOKEN_QUANT_ASYMM：per_token动态非对称量化，未实现。
- UNQUANT：不量化，浮点输出，未实现。


#### 输入

| 所属模块功能 | 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- | --- |
| rmsNormQuant\_0 | input | [tokenNum, hiddenSize] | float16/bf16 | ND | 必选。 |
| rmsNormQuant\_0 | gamma0 | [hiddenSize] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rmsNormQuant\_0 | beta0 | [hiddenSize] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rmsNormQuant\_0 | quantScale0 | [1] | float16/bf16 | ND | 必选，支持传入空tensor。仅在quantMode为0时传入，数据类型与input一致。 |
| rmsNormQuant\_0 | quantOffset0 | [1] | int8 | ND | 必选，支持传入空tensor。仅在quantMode为0时传入。 |
| matmul\_0 | wdqkv | [2112, hiddenSize] | int8/bf16 | NZ | 必选。 数据类型与input一致且为bf16时不开启rmsNormQuant。 |
| matmul\_0 | deScale0 | [2112] | int64/float | ND | 必选。input为float16时为int64，input为bf16时为float。 |
| matmul\_0 | bias0 | [2112] | int32 | ND | 必选，支持传入空tensor。quantMode为1、3时不传入。 |
| rmsNormQuant\_1 | gamma1 | [1536] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rmsNormQuant\_1 | beta1 | [1536] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rmsNormQuant\_1 | quantScale1 | [1] | float16/bf16 | ND | 必选，支持传入空tensor。仅在quantMode为0时传入，数据类型与input一致。 |
| rmsNormQuant\_1 | quantOffset1 | [1] | int8 | ND | 必选，支持传入空tensor。仅在quantMode为0时传入。 |
| matmul\_1 | wuq | [headNum \* 192, 1536] | int8/bf16 | NZ | 必选。数据类型与input一致且为bf16时不开启rmsNormQuant。 |
| matmul\_1 | deScale1 | [headNum \* 192] | int64/float | ND | 必选，input为float16时为int64，input为bf16时为float。 |
| matmul\_1 | bias1 | [headNum \* 192] | int32 | ND | 必选，支持传入空tensor。quantMode为1、3时不传入。 |
| rmsNorm | gamma2 | [512] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rope | cos | [tokenNum,64] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| rope | sin | [tokenNum,64] | float16/bf16 | ND | 必选。数据类型与input一致。 |
| matmulEin | wuk | ND：[headNum, 128, 512] NZ：[headNum, 32, 128, 16] | float16/bf16 | ND/NZ | 必选。数据类型与input一致。 |
| reshapeAndCache | kvCache | cacheMode为0： [blockNum,blockSize, 1, 576] cacheMode为1： [blockNum,blockSize, 1, 512] cacheMode为2： [blockNum, 1\*512/32, blockSize, 32] cacheMode为3： [blockNum, 1\*512/16, blockSize, 16] | float16/bf16/int8 | ND/NZ | 必选。数据类型与input一致。 cacheMode为1时，tensor的shape为拆分情况。 cacheMode为2时格式为NZ，类型为int8。 cacheMode为3时，格式为NZ。 |
| reshapeAndCache | kvCacheRope | cacheMode为1： [blockNum,blockSize, 1, 64] cacheMode为2或3： [blockNum, 1\*64 / 16 , blockSize, 16] | float16/bf16 | ND/NZ | 必选，支持传入空tensor。 cacheMode不为0时传入，数据类型与input一致。 cacheMode为2或3时，格式为NZ。 |
| reshapeAndCache | slotmapping | [tokenNum] | int32 | ND | 必选。 |
| quant | ctkvScale | [1] | float16/bf16 | ND | 必选，支持传入空tensor。cacheMode为2时传入，数据类型与input一致。 |
| quant | qNopeScale | [headNum] | float16/bf16 | ND | 必选，支持传入空tensor。cacheMode为2时传入，数据类型与input一致。 |


#### 输出

可选输出tensor不能使用空tensor占位。

| 所属模块功能 | 参数 | 维度 | 数据类型 | 格式 | 描述 |
| --- | --- | --- | --- | --- | --- |
| 输出数据 | qOut0 | cacheMode为0： [tokenNum,headNum, 576] cacheMode为1或2： [tokenNum,headNum, 512] | float16/bf16/int8 | ND | 输出tensor。数据类型与input一致。 cacheMode为2时数据类型为int8。 |
| 输出数据 | kvCacheOut0 | cacheMode为0： [blockNum, blockSize, 1, 576] cacheMode为1： [blockNum, blockSize, 1, 512] cacheMode为2： [blockNum, 1\*512/32, blockSize, 32] cacheMode为3： [blockNum, 1\*512/16, blockSize, 16] | float16/bf16/int8 | ND/NZ | 输出tensor。数据类型与input一致。 cacheMode为2时数据类型为int8，格式为NZ。 cacheMode为3时，格式为NZ。 |
| 输出数据 | qOut1 | [tokenNum, headNum, 64] | float16/bf16 | ND | cacheMode不为0时输出此tensor。数据类型与input一致。 |
| 输出数据 | kvCacheOut1 | cacheMode为1： [blockNum, blockSize, 1, 64] cacheMode为2： [blockNum, 1\*64/16, blockSize, 16] | float16/bf16 | ND/NZ | cacheMode不为0时输出此tensor。数据类型与input一致。 cacheMode为2时数据格式为NZ。 cacheMode为3时，格式为NZ。 |


#### 约束说明

- tokenNum<=1024。

- blockSize<=128 或 blockSize=256。
- cacheMode为2或3时，blockSize=128。
- hiddenSize取值：[2048, 8192]，参数需要256对齐。
- 不开启rmsNormQuant时，hiddenSize只能等于6144，且仅支持cacheMode=0、1。
- 不开启rmsNormQuant或hiddenSize不等于7168时：

  - 1<= tokenNum<=256。
  - headNum取值范围：[16, 32, 64, 128]。
